package models import ( "context" "errors" "io" "ragflow/internal/common" "testing" "github.com/cloudwego/eino/schema" ) func TestEinoChatModelStreamFiltersDoneSentinel(t *testing.T) { modelName := "chat" driver := &streamSentinelDriver{captureToolDriver: &captureToolDriver{}} base := NewChatModel(driver, &modelName, &APIConfig{}) model := NewEinoChatModel(base, &ChatConfig{}) stream, err := model.Stream(context.Background(), []*schema.Message{schema.UserMessage("hello")}) if err != nil { t.Fatalf("Stream: %v", err) } var messages []string for { msg, recvErr := stream.Recv() if errors.Is(recvErr, io.EOF) { break } if recvErr != nil { t.Fatalf("stream.Recv: %v", recvErr) } if msg != nil { messages = append(messages, msg.Content) } } if len(messages) != 2 || messages[0] != "answer" || messages[1] != "DONE!" { t.Fatalf("stream messages = %#v, want [answer DONE!]", messages) } } func TestEinoChatModelGenerateSendsBoundTools(t *testing.T) { apiKey := "key" modelName := "chat" driver := &captureToolDriver{ resp: &ChatResponse{ ToolCalls: []map[string]interface{}{ { "id": "call-1", "type": "function", "function": map[string]interface{}{ "name": "search_my_dateset", "arguments": `{"query":"hello"}`, }, }, }, }, } base := NewChatModel(driver, &modelName, &APIConfig{ApiKey: &apiKey}) model := NewEinoChatModel(base, nil) 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) } msg, err := bound.Generate(context.Background(), []*schema.Message{schema.UserMessage("hello")}) if err != nil { t.Fatalf("Generate: %v", err) } if driver.lastConfig == nil || driver.lastConfig.Tools == nil { t.Fatal("Generate did not send tools to driver") } tools, ok := driver.lastConfig.Tools.([]map[string]any) if !ok || len(tools) != 1 { t.Fatalf("driver tools = %#v, want one OpenAI-style tool", driver.lastConfig.Tools) } fn, _ := tools[0]["function"].(map[string]any) if fn["name"] != "search_my_dateset" { t.Fatalf("tool function name = %#v, want search_my_dateset", fn["name"]) } if driver.lastConfig.ToolChoice == nil || *driver.lastConfig.ToolChoice != "auto" { t.Fatalf("ToolChoice = %#v, want auto", driver.lastConfig.ToolChoice) } if len(msg.ToolCalls) != 1 { t.Fatalf("msg.ToolCalls len = %d, want 1", len(msg.ToolCalls)) } if msg.ToolCalls[0].Function.Name != "search_my_dateset" || msg.ToolCalls[0].Function.Arguments != `{"query":"hello"}` { t.Fatalf("tool call = %#v, want search_my_dateset query call", msg.ToolCalls[0]) } } func TestEinoChatModelStreamWithToolsYieldsToolCalls(t *testing.T) { apiKey := "key" modelName := "chat" driver := &captureToolDriver{ resp: &ChatResponse{ ToolCalls: []map[string]interface{}{ { "id": "call-1", "type": "function", "function": map[string]interface{}{ "name": "search_my_dateset", "arguments": `{"query":"hello"}`, }, }, }, }, } base := NewChatModel(driver, &modelName, &APIConfig{ApiKey: &apiKey}) model := NewEinoChatModel(base, nil) 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")}) if err != nil { t.Fatalf("Stream: %v", err) } msg, err := stream.Recv() if err != nil { t.Fatalf("stream.Recv: %v", err) } if msg == nil { t.Fatal("stream ended before yielding message") } 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}, } base := NewChatModel(driver, &modelName, &APIConfig{ApiKey: &apiKey}) model := NewEinoChatModel(base, &ChatConfig{}) 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) } } func TestToInternalMessagesPreservesToolMessages(t *testing.T) { internal := toInternalMessages([]*schema.Message{ { Role: schema.Assistant, ToolCalls: []schema.ToolCall{{ ID: "call-1", Type: "function", Function: schema.FunctionCall{ Name: "search_my_dateset", Arguments: `{"query":"hello"}`, }, }}, }, { Role: schema.Tool, Content: `{"formalized_content":"answer"}`, ToolCallID: "call-1", }, }) if len(internal) != 2 { t.Fatalf("len(internal) = %d, want 2", len(internal)) } if len(internal[0].ToolCalls) != 1 { t.Fatalf("assistant ToolCalls = %#v, want one tool call", internal[0].ToolCalls) } if internal[1].ToolCallID != "call-1" || internal[1].Role != "tool" { t.Fatalf("tool message = %#v, want tool role with call id", internal[1]) } } type captureToolDriver struct { resp *ChatResponse lastConfig *ChatConfig } type streamSentinelDriver struct { *captureToolDriver } func (d *streamSentinelDriver) ChatStreamlyWithSender(ctx context.Context, _ string, _ []Message, _ *APIConfig, _ *ChatConfig, _ *common.ModelUsage, sender func(*string, *string) error) error { answer := "answer" if err := sender(&answer, nil); err != nil { return err } visibleDone := "DONE!" if err := sender(&visibleDone, nil); err != nil { return err } done := "[DONE]" return sender(&done, nil) } func (d *captureToolDriver) NewInstance(baseURL map[string]string) ModelDriver { return d } func (d *captureToolDriver) Name() string { return "capture" } func (d *captureToolDriver) ChatWithMessages(ctx context.Context, _ string, _ []Message, _ *APIConfig, cfg *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) { d.lastConfig = cfg return d.resp, nil } func (d *captureToolDriver) ChatStreamlyWithSender(ctx context.Context, _ 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(ctx context.Context, _ *string, _ EmbedRequest, _ *APIConfig, _ *EmbeddingConfig, _ *common.ModelUsage) ([]EmbeddingData, error) { return nil, nil } func (d *captureToolDriver) Rerank(ctx context.Context, _ *string, _ RerankRequest, _ *APIConfig, _ *RerankConfig, _ *common.ModelUsage) (*RerankResponse, error) { return nil, nil } func (d *captureToolDriver) TranscribeAudio(ctx context.Context, _ *string, _ *string, _ *APIConfig, _ *ASRConfig, _ *common.ModelUsage) (*ASRResponse, error) { return nil, nil } func (d *captureToolDriver) TranscribeAudioWithSender(ctx context.Context, _ *string, _ *string, _ *APIConfig, _ *ASRConfig, _ *common.ModelUsage, _ func(*string, *string) error) error { return nil } func (d *captureToolDriver) AudioSpeech(ctx context.Context, _ *string, _ *string, _ *APIConfig, _ *TTSConfig, _ *common.ModelUsage) (*TTSResponse, error) { return nil, nil } func (d *captureToolDriver) AudioSpeechWithSender(ctx context.Context, _ *string, _ *string, _ *APIConfig, _ *TTSConfig, _ *common.ModelUsage, _ func(*string, *string) error) error { return nil } func (d *captureToolDriver) OCRFile(ctx context.Context, _ *string, _ []byte, _ *string, _ *APIConfig, _ *OCRConfig, _ *common.ModelUsage) (*OCRFileResponse, error) { return nil, nil } func (d *captureToolDriver) ParseFile(ctx context.Context, _ *string, _ []byte, _ *string, _ *APIConfig, _ *ParseFileConfig, _ *common.ModelUsage) (*ParseFileResponse, error) { return nil, nil } func (d *captureToolDriver) ListModels(ctx context.Context, _ *APIConfig) ([]ListModelResponse, error) { return nil, nil } func (d *captureToolDriver) Balance(ctx context.Context, _ *APIConfig) (map[string]interface{}, error) { return nil, nil } func (d *captureToolDriver) CheckConnection(ctx context.Context, _ *APIConfig) error { return nil } func (d *captureToolDriver) ListTasks(ctx context.Context, _ *APIConfig) ([]ListTaskStatus, error) { return nil, nil } func (d *captureToolDriver) ShowTask(ctx context.Context, _ string, _ *APIConfig) (*TaskResponse, error) { return nil, nil } // TestToInternalMessagesConvertsMultiModalContent guards the eino→driver // boundary: UserInputMultiContent must become OpenAI-style content blocks // ([]interface{} of {type:text} / {type:image_url}) on Message.Content, // otherwise image parts produced by the component layer are silently // dropped before the request reaches any driver. func TestToInternalMessagesConvertsMultiModalContent(t *testing.T) { uri := "data:image/png;base64,iVBORw0KGgo=" internal := toInternalMessages([]*schema.Message{ { Role: schema.User, UserInputMultiContent: []schema.MessageInputPart{ {Type: schema.ChatMessagePartTypeText, Text: "describe the image"}, {Type: schema.ChatMessagePartTypeImageURL, Image: &schema.MessageInputImage{ MessagePartCommon: schema.MessagePartCommon{URL: &uri}, }}, }, }, }) if len(internal) != 1 { t.Fatalf("len(internal) = %d, want 1", len(internal)) } blocks, ok := internal[0].Content.([]interface{}) if !ok { t.Fatalf("Content type = %T, want []interface{} content blocks", internal[0].Content) } if len(blocks) != 2 { t.Fatalf("len(blocks) = %d, want 2", len(blocks)) } textBlock, ok := blocks[0].(map[string]interface{}) if !ok || textBlock["type"] != "text" || textBlock["text"] != "describe the image" { t.Fatalf("text block = %#v, want {type:text, text:describe the image}", blocks[0]) } imageBlock, ok := blocks[1].(map[string]interface{}) if !ok || imageBlock["type"] != "image_url" { t.Fatalf("image block = %#v, want type image_url", blocks[1]) } imageURL, ok := imageBlock["image_url"].(map[string]interface{}) if !ok || imageURL["url"] != uri { t.Fatalf("image_url = %#v, want url %q", imageBlock["image_url"], uri) } } // TestToInternalMessagesReassemblesBase64Image: parts that carry Base64Data // instead of a URL are reassembled into a data URI. func TestToInternalMessagesReassemblesBase64Image(t *testing.T) { b64 := "aGVsbG8=" internal := toInternalMessages([]*schema.Message{ { Role: schema.User, UserInputMultiContent: []schema.MessageInputPart{ {Type: schema.ChatMessagePartTypeImageURL, Image: &schema.MessageInputImage{ MessagePartCommon: schema.MessagePartCommon{ Base64Data: &b64, MIMEType: "image/jpeg", }, }}, }, }, }) blocks, ok := internal[0].Content.([]interface{}) if !ok || len(blocks) != 1 { t.Fatalf("Content = %#v, want one content block", internal[0].Content) } imageBlock, ok := blocks[0].(map[string]interface{}) if !ok { t.Fatalf("block = %#v, want map", blocks[0]) } imageURL, ok := imageBlock["image_url"].(map[string]interface{}) if !ok || imageURL["url"] != "data:image/jpeg;base64,aGVsbG8=" { t.Fatalf("image_url = %#v, want reassembled data URI", imageBlock["image_url"]) } } // TestToInternalMessagesUnsupportedPartsFallBackToString: when every part is // of an unsupported type, Content stays the plain string. func TestToInternalMessagesUnsupportedPartsFallBackToString(t *testing.T) { internal := toInternalMessages([]*schema.Message{ { Role: schema.User, Content: "plain", UserInputMultiContent: []schema.MessageInputPart{ {Type: schema.ChatMessagePartTypeAudioURL}, }, }, }) if content, ok := internal[0].Content.(string); !ok || content != "plain" { t.Fatalf("Content = %#v, want string %q", internal[0].Content, "plain") } }