mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-16 05:26:06 +08:00
396 lines
13 KiB
Go
396 lines
13 KiB
Go
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")
|
|
}
|
|
}
|