Files
ragflow/internal/entity/models/llm_test.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")
}
}