fix: unable to got SEE response when use tools (#17076)

### Summary

As title
This commit is contained in:
Haruko386
2026-07-20 10:50:30 +08:00
committed by GitHub
parent 679d0d2ad2
commit a9db04ab23
2 changed files with 81 additions and 9 deletions

View File

@@ -276,14 +276,11 @@ func (m *EinoChatModel) Stream(ctx context.Context, msgs []*schema.Message, opts
if m.inner.ModelName == nil {
return nil, fmt.Errorf("models: EinoChatModel: nil model name")
}
if len(m.tools) > 0 {
msg, err := m.Generate(ctx, msgs, opts...)
if err != nil {
return nil, err
}
return schema.StreamReaderFromArray([]*schema.Message{msg}), nil
internalMessage := toInternalMessages(msgs)
chatCfg, err := m.chatConfigForGenerate()
if err != nil {
return nil, err
}
internal := toInternalMessages(msgs)
sr, sw := schema.Pipe[*schema.Message](1)
var sendMu sync.Mutex
@@ -316,8 +313,16 @@ func (m *EinoChatModel) Stream(ctx context.Context, msgs []*schema.Message, opts
}
go func() {
defer sw.Close()
if err := m.inner.ModelDriver.ChatStreamlyWithSender(*m.inner.ModelName, internal, m.inner.APIConfig, m.chatCfg, nil, sender); err != nil {
if err := m.inner.ModelDriver.ChatStreamlyWithSender(*m.inner.ModelName, internalMessage, m.inner.APIConfig, chatCfg, nil, sender); err != nil {
_ = sw.Send(nil, err)
return
}
if chatCfg != nil && chatCfg.ToolCallsResult != nil && len(*chatCfg.ToolCallsResult) > 0 {
msg := &schema.Message{
Role: schema.Assistant,
ToolCalls: toolCallsFromInternal(*chatCfg.ToolCallsResult),
}
_ = sw.Send(msg, nil)
}
}()
return sr, nil

View File

@@ -154,6 +154,61 @@ func TestEinoChatModelStreamWithToolsYieldsToolCalls(t *testing.T) {
if len(msg.ToolCalls) != 1 || msg.ToolCalls[0].Function.Name != "search_my_dateset" {
t.Fatalf("stream message tool calls = %#v, want search_my_dateset", msg.ToolCalls)
}
if driver.lastConfig == nil || driver.lastConfig.Tools == nil {
t.Fatal("Stream did not send tools to driver")
}
}
func TestEinoChatModelStreamWithToolsStreamsFinalAnswer(t *testing.T) {
apiKey := "key"
modelName := "chat"
answer := "streamed answer"
driver := &captureToolDriver{
resp: &ChatResponse{Answer: &answer},
}
var callbacks []string
base := NewChatModel(driver, &modelName, &APIConfig{ApiKey: &apiKey})
model := NewEinoChatModel(base, &ChatConfig{
StreamCallback: func(content, _ string) {
if content != "" {
callbacks = append(callbacks, content)
}
},
})
bound, err := model.WithTools([]*schema.ToolInfo{
{
Name: "search_my_dateset",
Desc: "Search datasets.",
ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{
"query": {Type: schema.String, Required: true},
}),
},
})
if err != nil {
t.Fatalf("WithTools: %v", err)
}
stream, err := bound.Stream(context.Background(), []*schema.Message{
schema.UserMessage("hello"),
{
Role: schema.Tool,
Content: `{"formalized_content":"hit"}`,
ToolCallID: "call-1",
},
})
if err != nil {
t.Fatalf("Stream: %v", err)
}
msg, err := stream.Recv()
if err != nil {
t.Fatalf("stream.Recv: %v", err)
}
if msg == nil || msg.Content != answer {
t.Fatalf("stream message = %#v, want final answer content", msg)
}
if len(callbacks) != 1 || callbacks[0] != answer {
t.Fatalf("callbacks = %#v, want streamed answer", callbacks)
}
}
func TestToInternalMessagesPreservesToolMessages(t *testing.T) {
@@ -214,7 +269,19 @@ func (d *captureToolDriver) ChatWithMessages(_ string, _ []Message, _ *APIConfig
d.lastConfig = cfg
return d.resp, nil
}
func (d *captureToolDriver) ChatStreamlyWithSender(_ string, _ []Message, _ *APIConfig, _ *ChatConfig, _ *common.ModelUsage, _ func(*string, *string) error) error {
func (d *captureToolDriver) ChatStreamlyWithSender(_ string, _ []Message, _ *APIConfig, cfg *ChatConfig, _ *common.ModelUsage, sender func(*string, *string) error) error {
d.lastConfig = cfg
if d.resp == nil {
return nil
}
if cfg != nil && len(d.resp.ToolCalls) > 0 {
tcs := append([]map[string]interface{}(nil), d.resp.ToolCalls...)
cfg.ToolCallsResult = &tcs
return nil
}
if d.resp.Answer != nil {
return sender(d.resp.Answer, d.resp.ReasonContent)
}
return nil
}
func (d *captureToolDriver) Embed(_ *string, _ []string, _ *APIConfig, _ *EmbeddingConfig, _ *common.ModelUsage) ([]EmbeddingData, error) {