diff --git a/internal/entity/models/llm.go b/internal/entity/models/llm.go index cfe163311c..8194fd54c1 100644 --- a/internal/entity/models/llm.go +++ b/internal/entity/models/llm.go @@ -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 diff --git a/internal/entity/models/llm_test.go b/internal/entity/models/llm_test.go index 147f05f8c4..5bd22d87d4 100644 --- a/internal/entity/models/llm_test.go +++ b/internal/entity/models/llm_test.go @@ -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) {