From 17d71e5d79207cad6b0e3cec72114440a2e0f45e Mon Sep 17 00:00:00 2001 From: Jin Hai Date: Sat, 9 May 2026 17:41:54 +0800 Subject: [PATCH] Go CLI: embed and rerank (#14735) ### What problem does this PR solve? ``` RAGFlow(user)> embed text 'what is rag' 'who are you' with 'embedding-3@test@zhipu-ai' dimension 16; +-----------+-------+ | dimension | index | +-----------+-------+ | 16 | 0 | | 16 | 1 | +-----------+-------+ RAGFlow(user)> rerank query 'what is rag' document 'rag is retrieval augment generation' 'rag need llm' 'famous rag project includes ragflow' with 'rerank@test@zhipu-ai' top 2; +-------+-----------------+ | index | relevance_score | +-------+-----------------+ | 0 | 1 | | 2 | 0.99999976 | +-------+-----------------+ ``` ### Type of change - [x] New Feature (non-breaking change which adds functionality) Signed-off-by: Jin Hai --- conf/models/zhipu-ai.json | 4 +- internal/cli/client.go | 4 + internal/cli/lexer.go | 10 ++ internal/cli/parser.go | 64 +++---- internal/cli/types.go | 5 + internal/cli/user_command.go | 148 +++++++++++++++- internal/cli/user_parser.go | 120 +++++++++++++ internal/common/float.go | 40 +++++ internal/entity/models/aliyun.go | 38 ++--- internal/entity/models/deepseek.go | 4 +- internal/entity/models/dummy.go | 4 +- internal/entity/models/gitee.go | 38 ++--- internal/entity/models/google.go | 4 +- internal/entity/models/huggingface.go | 2 +- internal/entity/models/lmstudio.go | 2 +- internal/entity/models/minimax.go | 4 +- internal/entity/models/moonshot.go | 4 +- internal/entity/models/nvidia.go | 2 +- internal/entity/models/ollama.go | 2 +- internal/entity/models/openai.go | 4 +- internal/entity/models/openrouter.go | 27 +-- internal/entity/models/siliconflow.go | 64 ++++--- internal/entity/models/types.go | 30 +++- internal/entity/models/vllm.go | 4 +- internal/entity/models/volcengine.go | 4 +- internal/entity/models/xai.go | 4 +- internal/entity/models/zhipu-ai.go | 51 ++++-- internal/handler/providers.go | 153 +++++++++++++++++ internal/router/router.go | 2 + internal/service/model_service.go | 232 ++++++++++++++++++++++++++ internal/service/nlp/reranker.go | 14 +- 31 files changed, 919 insertions(+), 169 deletions(-) create mode 100644 internal/common/float.go diff --git a/conf/models/zhipu-ai.json b/conf/models/zhipu-ai.json index 52f4a8396a..d1bbac649f 100644 --- a/conf/models/zhipu-ai.json +++ b/conf/models/zhipu-ai.json @@ -242,7 +242,7 @@ ] }, { - "name": "glm-asr", + "name": "glm-asr-2512", "max_tokens": 4096, "model_types": [ "asr" @@ -261,7 +261,7 @@ ] }, { - "name": "glm-rerank", + "name": "rerank", "model_types": [ "rerank" ] diff --git a/internal/cli/client.go b/internal/cli/client.go index 2a0a013799..2bd50cb695 100644 --- a/internal/cli/client.go +++ b/internal/cli/client.go @@ -263,6 +263,10 @@ func (c *RAGFlowClient) ExecuteUserCommand(cmd *Command) (ResponseIf, error) { return c.ChatToModel(cmd) case "think_chat_to_model": return c.ChatToModel(cmd) + case "embed_user_text": + return c.EmbedUserText(cmd) + case "rarank_user_document": + return c.RerankUserDocument(cmd) case "check_provider_connection": return c.CheckProviderConnection(cmd) case "use_model": diff --git a/internal/cli/lexer.go b/internal/cli/lexer.go index 59c23646ee..5f2aadea14 100644 --- a/internal/cli/lexer.go +++ b/internal/cli/lexer.go @@ -363,6 +363,16 @@ func (l *Lexer) lookupIdent(ident string) Token { return Token{Type: TokenASR, Value: ident} case "TTS": return Token{Type: TokenTTS, Value: ident} + case "EMBED": + return Token{Type: TokenEmbed, Value: ident} + case "TEXT": + return Token{Type: TokenText, Value: ident} + case "QUERY": + return Token{Type: TokenQuery, Value: ident} + case "TOP": + return Token{Type: TokenTop, Value: ident} + case "DIMENSION": + return Token{Type: TokenDimension, Value: ident} case "OCR": return Token{Type: TokenOCR, Value: ident} case "ASYNC": diff --git a/internal/cli/parser.go b/internal/cli/parser.go index 92908f2ea9..e373c5a874 100644 --- a/internal/cli/parser.go +++ b/internal/cli/parser.go @@ -197,6 +197,10 @@ func (p *Parser) parseUserCommand() (*Command, error) { return p.parseChatCommand() case TokenThink: return p.parseThinkCommand() + case TokenEmbed: + return p.parseEmbedCommand() + case TokenRerank: + return p.parseRerankCommand() case TokenCheck: return p.parseCheckCommand() case TokenLS: @@ -495,43 +499,43 @@ func (p *Parser) parseCESearchCommand() (*Command, error) { p.curToken.Type == TokenChats || p.curToken.Type == TokenDatasets { path = path + "/" + p.curToken.Value p.nextToken() - } else if p.curToken.Type == TokenNumber { - // Handle version numbers like 1.0.0 (parsed as number . number . number) - // OR filenames starting with numbers like 3_list_compressors.pdf - numberPart := p.curToken.Value - p.nextToken() - // Continue reading .number parts (version number format) - if p.curToken.Type == TokenIllegal && p.curToken.Value == "." { - versionPart := numberPart - for p.curToken.Type == TokenIllegal && p.curToken.Value == "." { - p.nextToken() // consume . - if p.curToken.Type == TokenNumber { - versionPart = versionPart + "." + p.curToken.Value - p.nextToken() - } else { - break + } else if p.curToken.Type == TokenNumber { + // Handle version numbers like 1.0.0 (parsed as number . number . number) + // OR filenames starting with numbers like 3_list_compressors.pdf + numberPart := p.curToken.Value + p.nextToken() + // Continue reading .number parts (version number format) + if p.curToken.Type == TokenIllegal && p.curToken.Value == "." { + versionPart := numberPart + for p.curToken.Type == TokenIllegal && p.curToken.Value == "." { + p.nextToken() // consume . + if p.curToken.Type == TokenNumber { + versionPart = versionPart + "." + p.curToken.Value + p.nextToken() + } else { + break + } } + path = path + "/" + versionPart + } else if p.curToken.Type == TokenIdentifier { + // Filename starting with number: 3_list_compressors.pdf + path = path + "/" + numberPart + p.curToken.Value + p.nextToken() + } else { + // Just a number + path = path + "/" + numberPart } - path = path + "/" + versionPart - } else if p.curToken.Type == TokenIdentifier { - // Filename starting with number: 3_list_compressors.pdf - path = path + "/" + numberPart + p.curToken.Value + } else if p.curToken.Type == TokenQuotedString { + path = path + "/" + strings.Trim(p.curToken.Value, "\"'") p.nextToken() } else { - // Just a number - path = path + "/" + numberPart + // Trailing slash, just append it + path = path + "/" + break } - } else if p.curToken.Type == TokenQuotedString { - path = path + "/" + strings.Trim(p.curToken.Value, "\"'") - p.nextToken() - } else { - // Trailing slash, just append it - path = path + "/" - break } - } - cmd.Params["path"] = path + cmd.Params["path"] = path } else { cmd.Params["path"] = "." } diff --git a/internal/cli/types.go b/internal/cli/types.go index 9a373df87a..a30f26c6ad 100644 --- a/internal/cli/types.go +++ b/internal/cli/types.go @@ -102,6 +102,11 @@ const ( TokenASR TokenTTS TokenOCR + TokenEmbed + TokenText + TokenQuery + TokenTop + TokenDimension TokenAsync TokenSync TokenBenchmark diff --git a/internal/cli/user_command.go b/internal/cli/user_command.go index 6dbf84be25..a8394e40a6 100644 --- a/internal/cli/user_command.go +++ b/internal/cli/user_command.go @@ -1572,7 +1572,6 @@ func (c *RAGFlowClient) ChatToModel(cmd *Command) (ResponseIf, error) { "text": message, }) } - } images, ok := cmd.Params["images"].([]string) @@ -1783,6 +1782,146 @@ func (c *RAGFlowClient) ChatToModel(cmd *Command) (ResponseIf, error) { return &result, nil } +func (c *RAGFlowClient) EmbedUserText(cmd *Command) (ResponseIf, error) { + if c.HTTPClient.APIToken == "" && c.HTTPClient.LoginToken == "" { + return nil, fmt.Errorf("API token not set. Please login first") + } + + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + var providerName, instanceName, modelName string + + // Check if composite_model_name is provided in command + if compositeModelName, ok := cmd.Params["composite_model_name"].(string); ok && compositeModelName != "" { + names := strings.Split(compositeModelName, "@") + if len(names) != 3 { + return nil, fmt.Errorf("model name must be in format 'model@instance@provider'") + } + providerName = names[2] + instanceName = names[1] + modelName = names[0] + } else if c.CurrentModel != nil { + // Use current model if set + providerName = c.CurrentModel.Provider + instanceName = c.CurrentModel.Instance + modelName = c.CurrentModel.Model + } else { + return nil, fmt.Errorf("model name not provided and no current model set. Use 'use model' command first") + } + + texts, ok := cmd.Params["texts"].([]string) + if !ok { + return nil, fmt.Errorf("texts not provided") + } + + dimension, ok := cmd.Params["dimension"].(int) + if !ok { + dimension = 0 + } + + payload := map[string]interface{}{ + "provider_name": providerName, + "instance_name": instanceName, + "model_name": modelName, + "texts": texts, + "dimension": dimension, + } + + url := "/embeddings" + + resp, err := c.HTTPClient.Request("POST", url, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to embed text: %w", err) + } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to embed text: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + } + var result CommonResponse + if err = json.Unmarshal(resp.Body, &result); err != nil { + return nil, fmt.Errorf("embed text failed: invalid JSON (%w)", err) + } + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + result.Duration = resp.Duration + return &result, nil +} + +func (c *RAGFlowClient) RerankUserDocument(cmd *Command) (ResponseIf, error) { + if c.HTTPClient.APIToken == "" && c.HTTPClient.LoginToken == "" { + return nil, fmt.Errorf("API token not set. Please login first") + } + + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + var providerName, instanceName, modelName string + + // Check if composite_model_name is provided in command + if compositeModelName, ok := cmd.Params["composite_model_name"].(string); ok && compositeModelName != "" { + names := strings.Split(compositeModelName, "@") + if len(names) != 3 { + return nil, fmt.Errorf("model name must be in format 'model@instance@provider'") + } + providerName = names[2] + instanceName = names[1] + modelName = names[0] + } else if c.CurrentModel != nil { + // Use current model if set + providerName = c.CurrentModel.Provider + instanceName = c.CurrentModel.Instance + modelName = c.CurrentModel.Model + } else { + return nil, fmt.Errorf("model name not provided and no current model set. Use 'use model' command first") + } + + query, ok := cmd.Params["query"].(string) + if !ok { + return nil, fmt.Errorf("query not provided") + } + + documents, ok := cmd.Params["documents"].([]string) + if !ok { + return nil, fmt.Errorf("documents not provided") + } + + topN, ok := cmd.Params["top_n"].(int) + if !ok { + return nil, fmt.Errorf("top n not provided") + } + + payload := map[string]interface{}{ + "provider_name": providerName, + "instance_name": instanceName, + "model_name": modelName, + "query": query, + "documents": documents, + "top_n": topN, + } + + url := "/rerank" + + resp, err := c.HTTPClient.Request("POST", url, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to rerank document: %w", err) + } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to rerank document: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + } + var result CommonResponse + if err = json.Unmarshal(resp.Body, &result); err != nil { + return nil, fmt.Errorf("rerank document failed: invalid JSON (%w)", err) + } + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + result.Duration = resp.Duration + return &result, nil +} + func (c *RAGFlowClient) CheckProviderConnection(cmd *Command) (ResponseIf, error) { if c.HTTPClient.APIToken == "" && c.HTTPClient.LoginToken == "" { return nil, fmt.Errorf("API token not set. Please login first") @@ -1820,7 +1959,6 @@ func (c *RAGFlowClient) CheckProviderConnection(cmd *Command) (ResponseIf, error } result.Duration = resp.Duration return &result, nil - } // UseModel sets the current model for chat @@ -1928,14 +2066,14 @@ func (c *RAGFlowClient) AddCustomModel(cmd *Command) (ResponseIf, error) { resp, err := c.HTTPClient.Request("POST", url, "web", nil, payload) if err != nil { - return nil, fmt.Errorf("failed to check provider connection: %w", err) + return nil, fmt.Errorf("failed to add custom model: %w", err) } if resp.StatusCode != 200 { - return nil, fmt.Errorf("failed to check provider connection: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + return nil, fmt.Errorf("failed to add custom model: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) } var result SimpleResponse if err = json.Unmarshal(resp.Body, &result); err != nil { - return nil, fmt.Errorf("check provider connection failed: invalid JSON (%w)", err) + return nil, fmt.Errorf("add custom model failed: invalid JSON (%w)", err) } if result.Code != 0 { return nil, fmt.Errorf("%s", result.Message) diff --git a/internal/cli/user_parser.go b/internal/cli/user_parser.go index ac6bbf358e..c49eeee11a 100644 --- a/internal/cli/user_parser.go +++ b/internal/cli/user_parser.go @@ -2603,6 +2603,126 @@ func (p *Parser) parseStreamCommand() (*Command, error) { return command, nil } +func (p *Parser) parseEmbedCommand() (*Command, error) { + p.nextToken() // consume EMBED + + if p.curToken.Type != TokenText { + return nil, fmt.Errorf("expected WITH after EMBED") + } + p.nextToken() // consume TEXT + + var texts []string + +textLoop: + for { + if p.curToken.Type != TokenQuotedString { + break textLoop + } + text, err := p.parseQuotedString() + if err != nil { + return nil, err + } + text = strings.TrimSpace(text) + texts = append(texts, text) + p.nextToken() + } + + if p.curToken.Type != TokenWith { + return nil, fmt.Errorf("expected WITH after EMBED") + } + p.nextToken() // consume WITH + + compositeModelName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + if p.curToken.Type != TokenDimension { + return nil, fmt.Errorf("expected DIMENSION") + } + p.nextToken() // consume WITH + + dimension, err := p.parseNumber() + if err != nil { + return nil, err + } + p.nextToken() + + cmd := NewCommand("embed_user_text") + cmd.Params["composite_model_name"] = compositeModelName + cmd.Params["texts"] = texts + cmd.Params["dimension"] = dimension + return cmd, nil +} + +func (p *Parser) parseRerankCommand() (*Command, error) { + p.nextToken() // consume RERANK + + if p.curToken.Type != TokenQuery { + return nil, fmt.Errorf("expected WITH after EMBED") + } + p.nextToken() // consume QUERY + + query, err := p.parseQuotedString() + if err != nil { + return nil, err + } + query = strings.TrimSpace(query) + p.nextToken() // consume query + + if p.curToken.Type != TokenDocument { + return nil, fmt.Errorf("expected DOCUMENT after query") + } + p.nextToken() // consume DOCUMENT + + var documents []string + +documentLoop: + for { + if p.curToken.Type != TokenQuotedString { + break documentLoop + } + var document string + document, err = p.parseQuotedString() + if err != nil { + return nil, err + } + document = strings.TrimSpace(document) + documents = append(documents, document) + p.nextToken() + } + + if p.curToken.Type != TokenWith { + return nil, fmt.Errorf("expected WITH after EMBED") + } + p.nextToken() // consume WITH + + compositeModelName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + if p.curToken.Type != TokenTop { + return nil, fmt.Errorf("expected TOP after model") + } + p.nextToken() + + topN, err := p.parseNumber() + if err != nil { + return nil, err + } + p.nextToken() + + cmd := NewCommand("rarank_user_document") + cmd.Params["composite_model_name"] = compositeModelName + cmd.Params["query"] = query + cmd.Params["documents"] = documents + cmd.Params["top_n"] = topN + return cmd, nil +} + func (p *Parser) parseCheckCommand() (*Command, error) { p.nextToken() // consume CHECK diff --git a/internal/common/float.go b/internal/common/float.go new file mode 100644 index 0000000000..b3dca37784 --- /dev/null +++ b/internal/common/float.go @@ -0,0 +1,40 @@ +// +// Copyright 2026 The InfiniFlow Authors. All Rights Reserved. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// + +package common + +const epsilon32 = 1e-6 +const epsilon64 = 1e-9 + +func Float64IsZero(f float64) bool { + if f < 0 && f >= -epsilon64 { + return true + } + if f > 0 && f <= epsilon64 { + return true + } + return false +} + +func Float32IsNotZero(f float32) bool { + if f < 0 && f >= -epsilon32 { + return true + } + if f > 0 && f <= epsilon32 { + return true + } + return false +} diff --git a/internal/entity/models/aliyun.go b/internal/entity/models/aliyun.go index 8fa546e0e7..a1ddd6dddb 100644 --- a/internal/entity/models/aliyun.go +++ b/internal/entity/models/aliyun.go @@ -473,9 +473,9 @@ type aliyunRerankResponse struct { } `json:"results"` } -func (z *AliyunModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - if len(texts) == 0 { - return []float64{}, nil +func (z *AliyunModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + if len(documents) == 0 { + return &RerankResponse{}, nil } if apiConfig == nil || apiConfig.ApiKey == nil || *apiConfig.ApiKey == "" { return nil, fmt.Errorf("api key is required") @@ -501,11 +501,16 @@ func (z *AliyunModel) Rerank(modelName *string, query string, texts []string, ap url := fmt.Sprintf("%s/%s", strings.TrimSuffix(baseURL, "/"), z.URLSuffix.Rerank) + var topN = rerankConfig.TopN + if rerankConfig.TopN == 0 { + topN = len(documents) + } + reqBody := aliyunRerankRequest{ Model: *modelName, Query: query, - Documents: texts, - TopN: len(texts), + Documents: documents, + TopN: topN, ReturnDocuments: false, } @@ -537,29 +542,12 @@ func (z *AliyunModel) Rerank(modelName *string, query string, texts []string, ap return nil, fmt.Errorf("Aliyun rerank API error: %s, body: %s", resp.Status, string(body)) } - var rerankResp aliyunRerankResponse - if err = json.Unmarshal(body, &rerankResp); err != nil { + var rerankResponse RerankResponse + if err = json.Unmarshal(body, &rerankResponse); err != nil { return nil, fmt.Errorf("failed to parse response: %w", err) } - scores := make([]float64, len(texts)) - seen := make([]bool, len(texts)) - for _, r := range rerankResp.Results { - if r.Index < 0 || r.Index >= len(texts) { - return nil, fmt.Errorf("aliyun rerank: result index %d out of range for %d documents", r.Index, len(texts)) - } - if seen[r.Index] { - return nil, fmt.Errorf("aliyun rerank: duplicate result index %d", r.Index) - } - scores[r.Index] = r.RelevanceScore - seen[r.Index] = true - } - - if len(rerankResp.Results) != len(texts) { - return nil, fmt.Errorf("aliyun rerank: expected %d results, got %d", len(texts), len(rerankResp.Results)) - } - - return scores, nil + return &rerankResponse, nil } type AliyunModelItem struct { diff --git a/internal/entity/models/deepseek.go b/internal/entity/models/deepseek.go index f1fd3116ac..dc06ebbfbd 100644 --- a/internal/entity/models/deepseek.go +++ b/internal/entity/models/deepseek.go @@ -580,7 +580,7 @@ func (z *DeepSeekModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts -func (z *DeepSeekModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *DeepSeekModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/dummy.go b/internal/entity/models/dummy.go index 124ba47309..ffc0f9f4b7 100644 --- a/internal/entity/models/dummy.go +++ b/internal/entity/models/dummy.go @@ -69,7 +69,7 @@ func (z *DummyModel) CheckConnection(apiConfig *APIConfig) error { return fmt.Errorf("no such method") } -// Rerank calculates similarity scores between query and texts -func (z *DummyModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *DummyModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/gitee.go b/internal/entity/models/gitee.go index 85d4635611..34d0425102 100644 --- a/internal/entity/models/gitee.go +++ b/internal/entity/models/gitee.go @@ -411,17 +411,10 @@ type giteeRerankRequest struct { ReturnDocuments bool `json:"return_documents"` } -type giteeRerankResponse struct { - Results []struct { - Index int `json:"index"` - RelevanceScore float64 `json:"relevance_score"` - } `json:"results"` -} - -// Rerank calculates similarity scores between query and texts -func (z *GiteeModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - if len(texts) == 0 { - return []float64{}, nil +// Rerank calculates similarity scores between query and documents +func (z *GiteeModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + if len(documents) == 0 { + return &RerankResponse{}, nil } if apiConfig == nil || apiConfig.ApiKey == nil || *apiConfig.ApiKey == "" { @@ -449,11 +442,16 @@ func (z *GiteeModel) Rerank(modelName *string, query string, texts []string, api url := fmt.Sprintf("%s/%s", strings.TrimSuffix(baseURL, "/"), z.URLSuffix.Rerank) + var topN = rerankConfig.TopN + if rerankConfig.TopN == 0 { + topN = len(documents) + } + reqBody := giteeRerankRequest{ Model: *modelName, Query: query, - Documents: texts, - TopN: len(texts), + Documents: documents, + TopN: topN, ReturnDocuments: false, } @@ -488,20 +486,12 @@ func (z *GiteeModel) Rerank(modelName *string, query string, texts []string, api return nil, fmt.Errorf("Gitee rerank API error: %s, body: %s", resp.Status, string(body)) } - var parsed giteeRerankResponse - if err = json.Unmarshal(body, &parsed); err != nil { + var rerankResponse RerankResponse + if err = json.Unmarshal(body, &rerankResponse); err != nil { return nil, fmt.Errorf("failed to parse response: %w", err) } - scores := make([]float64, len(texts)) - for _, r := range parsed.Results { - if r.Index < 0 || r.Index >= len(texts) { - return nil, fmt.Errorf("unexpected rerank index %d for %d inputs", r.Index, len(texts)) - } - scores[r.Index] = r.RelevanceScore - } - - return scores, nil + return &rerankResponse, nil } func (z *GiteeModel) ListModels(apiConfig *APIConfig) ([]string, error) { diff --git a/internal/entity/models/google.go b/internal/entity/models/google.go index d442b66399..b5679ac8da 100644 --- a/internal/entity/models/google.go +++ b/internal/entity/models/google.go @@ -248,7 +248,7 @@ func (z *GoogleModel) CheckConnection(apiConfig *APIConfig) error { return fmt.Errorf("no such method") } -// Rerank calculates similarity scores between query and texts -func (z *GoogleModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *GoogleModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/huggingface.go b/internal/entity/models/huggingface.go index 0c9e3ba5da..d1160d1c46 100644 --- a/internal/entity/models/huggingface.go +++ b/internal/entity/models/huggingface.go @@ -412,7 +412,7 @@ func (h *HuggingFaceModel) Encode(modelName *string, texts []string, apiConfig * return result, nil } -func (h *HuggingFaceModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (h *HuggingFaceModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("no such method") } diff --git a/internal/entity/models/lmstudio.go b/internal/entity/models/lmstudio.go index b9d1fee277..89a40e4685 100644 --- a/internal/entity/models/lmstudio.go +++ b/internal/entity/models/lmstudio.go @@ -365,7 +365,7 @@ func (l *LmStudioModel) Encode(modelName *string, texts []string, apiConfig *API return nil, fmt.Errorf("no such method") } -func (l *LmStudioModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (l *LmStudioModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("no such method") } diff --git a/internal/entity/models/minimax.go b/internal/entity/models/minimax.go index 04f5b1a02f..d40bfef4bd 100644 --- a/internal/entity/models/minimax.go +++ b/internal/entity/models/minimax.go @@ -443,7 +443,7 @@ func (z *MinimaxModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts -func (z *MinimaxModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *MinimaxModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/moonshot.go b/internal/entity/models/moonshot.go index 9d0de2c051..68af2fada8 100644 --- a/internal/entity/models/moonshot.go +++ b/internal/entity/models/moonshot.go @@ -483,7 +483,7 @@ func (z *MoonshotModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts -func (z *MoonshotModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *MoonshotModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/nvidia.go b/internal/entity/models/nvidia.go index 6a5f5907b9..4fd6a9b320 100644 --- a/internal/entity/models/nvidia.go +++ b/internal/entity/models/nvidia.go @@ -333,7 +333,7 @@ func (n NvidiaModel) Encode(modelName *string, texts []string, apiConfig *APICon return nil, fmt.Errorf("no such method") } -func (n NvidiaModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (n NvidiaModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("no such method") } diff --git a/internal/entity/models/ollama.go b/internal/entity/models/ollama.go index f2352bc6a8..4e8e42ad0d 100644 --- a/internal/entity/models/ollama.go +++ b/internal/entity/models/ollama.go @@ -363,7 +363,7 @@ func (o *OllamaModel) Encode(modelName *string, texts []string, apiConfig *APICo return nil, fmt.Errorf("no such method") } -func (o *OllamaModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (o *OllamaModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("no such method") } diff --git a/internal/entity/models/openai.go b/internal/entity/models/openai.go index f83d5810d4..1adbb35cbc 100644 --- a/internal/entity/models/openai.go +++ b/internal/entity/models/openai.go @@ -495,8 +495,8 @@ func (z *OpenAIModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts. OpenAI does +// Rerank calculates similarity scores between query and documents. OpenAI does // not expose a rerank API, so this is left unimplemented. -func (z *OpenAIModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (z *OpenAIModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/openrouter.go b/internal/entity/models/openrouter.go index b5ab500d11..505af9ee6a 100644 --- a/internal/entity/models/openrouter.go +++ b/internal/entity/models/openrouter.go @@ -470,9 +470,9 @@ type OpenRouterRerankResponse struct { } `json:"results"` } -func (o *OpenRouterModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - if len(texts) == 0 { - return []float64{}, nil +func (o *OpenRouterModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + if len(documents) == 0 { + return &RerankResponse{}, nil } var region = "default" @@ -480,11 +480,16 @@ func (o *OpenRouterModel) Rerank(modelName *string, query string, texts []string region = *apiConfig.Region } + var topN = rerankConfig.TopN + if rerankConfig.TopN == 0 { + topN = len(documents) + } + reqBody := OpenRouterRerankRequest{ Model: *modelName, Query: query, - Documents: texts, - TopN: len(texts), + Documents: documents, + TopN: topN, } jsonData, err := json.Marshal(reqBody) @@ -522,16 +527,16 @@ func (o *OpenRouterModel) Rerank(modelName *string, query string, texts []string return nil, fmt.Errorf("failed to decode response: %w", err) } - scores := make([]float64, len(texts)) - + var rerankResponse RerankResponse for _, result := range rerankResp.Results { - if result.Index >= 0 && - result.Index < len(texts) { - scores[result.Index] = result.RelevanceScore + rerankResult := RerankResult{ + Index: result.Index, + RelevanceScore: result.RelevanceScore, } + rerankResponse.Data = append(rerankResponse.Data, rerankResult) } - return scores, nil + return &rerankResponse, nil } func (o *OpenRouterModel) ListModels(apiConfig *APIConfig) ([]string, error) { diff --git a/internal/entity/models/siliconflow.go b/internal/entity/models/siliconflow.go index 61a300ce69..f3c658662c 100644 --- a/internal/entity/models/siliconflow.go +++ b/internal/entity/models/siliconflow.go @@ -72,14 +72,6 @@ type SiliconflowRerankRequest struct { OverlapTokens int `json:"overlap_tokens"` } -// SiliconflowRerankResponse represents SILICONFLOW rerank response -type SiliconflowRerankResponse struct { - Results []struct { - Index int `json:"index"` - RelevanceScore float64 `json:"relevance_score"` - } `json:"results"` -} - // ChatWithMessages sends multiple messages with roles and returns response func (z *SiliconflowModel) ChatWithMessages(modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig) (*ChatResponse, error) { if apiConfig == nil || apiConfig.ApiKey == nil || *apiConfig.ApiKey == "" { @@ -623,10 +615,36 @@ func (z *SiliconflowModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts -func (s *SiliconflowModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - if len(texts) == 0 { - return []float64{}, nil +// SiliconflowRerankResponse represents SILICONFLOW rerank response +type SiliconflowRerankResponse struct { + ID string `json:"id"` + Results []struct { + Index int `json:"index"` + Document struct { + Text string `json:"text"` + } `json:"document"` + RelevanceScore float64 `json:"relevance_score"` + } `json:"results"` + Meta struct { + Tokens struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + ImageTokens int `json:"image_tokens"` + } `json:"tokens"` + BilledUnits struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + ImageTokens int `json:"image_tokens"` + SearchUnits int `json:"search_units"` + Classifications int `json:"classifications"` + } `json:"billed_units"` + } `json:"meta"` +} + +// Rerank calculates similarity scores between query and documents +func (s *SiliconflowModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + if len(documents) == 0 { + return &RerankResponse{}, nil } var region = "default" @@ -642,8 +660,8 @@ func (s *SiliconflowModel) Rerank(modelName *string, query string, texts []strin reqBody := SiliconflowRerankRequest{ Model: *modelName, Query: query, - Documents: texts, - TopN: len(texts), + Documents: documents, + TopN: rerankConfig.TopN, ReturnDocuments: false, MaxChunksPerDoc: 1024, OverlapTokens: 80, @@ -679,17 +697,17 @@ func (s *SiliconflowModel) Rerank(modelName *string, query string, texts []strin body, _ := io.ReadAll(resp.Body) - var rerankResp SiliconflowRerankResponse - if err := json.Unmarshal(body, &rerankResp); err != nil { + var siliconflowRerankResp SiliconflowRerankResponse + if err = json.Unmarshal(body, &siliconflowRerankResp); err != nil { return nil, fmt.Errorf("failed to decode response: %w", err) } - scores := make([]float64, len(texts)) - for _, result := range rerankResp.Results { - if result.Index >= 0 && result.Index < len(texts) { - scores[result.Index] = result.RelevanceScore - } + var rerankResponse RerankResponse + for _, result := range siliconflowRerankResp.Results { + rerankResponse.Data = append(rerankResponse.Data, RerankResult{ + Index: result.Index, + RelevanceScore: result.RelevanceScore, + }) } - - return scores, nil + return &rerankResponse, nil } diff --git a/internal/entity/models/types.go b/internal/entity/models/types.go index 4833cf28f3..250e41bc51 100644 --- a/internal/entity/models/types.go +++ b/internal/entity/models/types.go @@ -25,7 +25,7 @@ type ModelDriver interface { // Encode encodes a list of texts into embeddings Encode(modelName *string, texts []string, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig) ([][]float64, error) // Rerank calculates similarity scores between query and texts - Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) + Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) // ListModels List supported models ListModels(apiConfig *APIConfig) ([]string, error) @@ -39,6 +39,25 @@ type ChatResponse struct { ReasonContent *string `json:"reason_content"` } +type EmbeddingResult struct { + Index int `json:"index"` + Dimension int `json:"dimension"` + //Embedding []float64 `json:"embedding"` +} + +type EmbeddingResponse struct { + Data []EmbeddingResult `json:"data"` +} + +type RerankResult struct { + Index int `json:"index"` + RelevanceScore float64 `json:"relevance_score"` +} + +type RerankResponse struct { + Data []RerankResult `json:"data"` +} + // URLSuffix represents the URL suffixes for different API endpoints type URLSuffix struct { Chat string `json:"chat"` @@ -72,6 +91,11 @@ type APIConfig struct { } type EmbeddingConfig struct { + Dimension int +} + +type RerankConfig struct { + TopN int } // EmbeddingModel wraps a ModelDriver with embedding-specific configuration @@ -109,8 +133,8 @@ func NewRerankModel(driver ModelDriver, modelName *string, apiConfig *APIConfig) } // Rerank calculates similarity between query and texts -func (r *RerankModel) Rerank(query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - return r.ModelDriver.Rerank(r.ModelName, query, texts, apiConfig) +func (r *RerankModel) Rerank(query string, texts []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + return r.ModelDriver.Rerank(r.ModelName, query, texts, apiConfig, rerankConfig) } // ChatModel wraps a ModelDriver with chat-specific configuration diff --git a/internal/entity/models/vllm.go b/internal/entity/models/vllm.go index b1ffe578fe..97ade07d1e 100644 --- a/internal/entity/models/vllm.go +++ b/internal/entity/models/vllm.go @@ -461,7 +461,7 @@ func (z *VllmModel) CheckConnection(apiConfig *APIConfig) error { return err } -// Rerank calculates similarity scores between query and texts -func (z *VllmModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *VllmModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/volcengine.go b/internal/entity/models/volcengine.go index 8b7ee8dab4..8b5670756d 100644 --- a/internal/entity/models/volcengine.go +++ b/internal/entity/models/volcengine.go @@ -490,8 +490,8 @@ func (z *VolcEngine) Encode(modelName *string, texts []string, apiConfig *APICon return embeddings, nil } -// Rerank calculates similarity scores between query and texts -func (z *VolcEngine) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +// Rerank calculates similarity scores between query and documents +func (z *VolcEngine) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/xai.go b/internal/entity/models/xai.go index afc6cc3dd3..96617320cf 100644 --- a/internal/entity/models/xai.go +++ b/internal/entity/models/xai.go @@ -487,8 +487,8 @@ func (z *XAIModel) CheckConnection(apiConfig *APIConfig) error { return nil } -// Rerank calculates similarity scores between query and texts. xAI does not +// Rerank calculates similarity scores between query and documents. xAI does not // expose a rerank API, so this is left unimplemented. -func (z *XAIModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { +func (z *XAIModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { return nil, fmt.Errorf("%s, Rerank not implemented", z.Name()) } diff --git a/internal/entity/models/zhipu-ai.go b/internal/entity/models/zhipu-ai.go index e0de7d8263..98bd5a7a52 100644 --- a/internal/entity/models/zhipu-ai.go +++ b/internal/entity/models/zhipu-ai.go @@ -374,9 +374,11 @@ func (z *ZhipuAIModel) Encode(modelName *string, texts []string, apiConfig *APIC embeddings := make([][]float64, len(texts)) for i, text := range texts { - reqBody := map[string]interface{}{ - "model": modelName, - "input": text, + reqBody := map[string]interface{}{} + reqBody["model"] = modelName + reqBody["input"] = text + if embeddingConfig.Dimension > 0 { + reqBody["dimensions"] = embeddingConfig.Dimension } jsonData, err := json.Marshal(reqBody) @@ -503,18 +505,26 @@ type zhipuRerankRequest struct { // zhipuRerankResponse is the response shape for the ZhipuAI rerank // endpoint. type zhipuRerankResponse struct { + Created int64 `json:"created"` + ID string `json:"id"` + RequestID string `json:"request_id"` + Usage struct { + CompletionTokens int `json:"completion_tokens"` + PromptTokens int `json:"prompt_tokens"` + TotalTokens int `json:"total_tokens"` + } `json:"usage"` Results []struct { Index int `json:"index"` RelevanceScore float64 `json:"relevance_score"` } `json:"results"` } -// Rerank calculates similarity scores between query and texts using +// Rerank calculates similarity scores between query and documents using // the ZhipuAI /rerank endpoint (e.g. glm-rerank). The result is one -// score per input text, in the same order the texts were given. -func (z *ZhipuAIModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, error) { - if len(texts) == 0 { - return []float64{}, nil +// score per input text, in the same order the documents were given. +func (z *ZhipuAIModel) Rerank(modelName *string, query string, documents []string, apiConfig *APIConfig, rerankConfig *RerankConfig) (*RerankResponse, error) { + if len(documents) == 0 { + return &RerankResponse{}, nil } if apiConfig == nil || apiConfig.ApiKey == nil || *apiConfig.ApiKey == "" { @@ -537,11 +547,16 @@ func (z *ZhipuAIModel) Rerank(modelName *string, query string, texts []string, a url := fmt.Sprintf("%s/%s", strings.TrimSuffix(baseURL, "/"), z.URLSuffix.Rerank) + var topN = rerankConfig.TopN + if rerankConfig.TopN == 0 { + topN = len(documents) + } + reqBody := zhipuRerankRequest{ Model: *modelName, Query: query, - Documents: texts, - TopN: len(texts), + Documents: documents, + TopN: topN, ReturnDocuments: false, } @@ -573,17 +588,19 @@ func (z *ZhipuAIModel) Rerank(modelName *string, query string, texts []string, a return nil, fmt.Errorf("ZhipuAI rerank API error: %s, body: %s", resp.Status, string(body)) } - var rerankResp zhipuRerankResponse - if err = json.Unmarshal(body, &rerankResp); err != nil { + var zhipuRerankResp zhipuRerankResponse + if err = json.Unmarshal(body, &zhipuRerankResp); err != nil { return nil, fmt.Errorf("failed to parse response: %w", err) } - scores := make([]float64, len(texts)) - for _, r := range rerankResp.Results { - if r.Index >= 0 && r.Index < len(texts) { - scores[r.Index] = r.RelevanceScore + var rerankResponse RerankResponse + for _, result := range zhipuRerankResp.Results { + rerankResult := RerankResult{ + Index: result.Index, + RelevanceScore: result.RelevanceScore, } + rerankResponse.Data = append(rerankResponse.Data, rerankResult) } - return scores, nil + return &rerankResponse, nil } diff --git a/internal/handler/providers.go b/internal/handler/providers.go index d90433cea5..758919f406 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -894,3 +894,156 @@ func (h *ProviderHandler) ChatToModel(c *gin.Context) { "answer": response.Answer, }) } + +type EmbedTextRequest struct { + ProviderName *string `json:"provider_name"` + InstanceName *string `json:"instance_name"` + ModelName *string `json:"model_name"` + Texts []string `json:"texts"` + Dimension int `json:"dimension"` +} + +func (h *ProviderHandler) EmbedText(c *gin.Context) { + var req EmbedTextRequest + if err := c.ShouldBindJSON(&req); err != nil { + println("JSON bind error: %v (type: %T)", err, err) + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": err.Error(), + }) + return + } + + if req.ProviderName == nil || *req.ProviderName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + if req.InstanceName == nil || *req.InstanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + if req.ModelName == nil || *req.ModelName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model name is required", + }) + return + } + + userID := c.GetString("user_id") + + apiConfig := models.APIConfig{ + ApiKey: nil, + Region: nil, + } + + embeddingConfig := models.EmbeddingConfig{ + Dimension: req.Dimension, + } + + // Non-stream response + var response *models.EmbeddingResponse + var errorCode common.ErrorCode + var err error + + response, errorCode, err = h.modelProviderService.EmbedText(*req.ProviderName, *req.InstanceName, *req.ModelName, userID, req.Texts, &apiConfig, &embeddingConfig) + + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "data": response.Data, + "message": "success", + }) +} + +type RerankDocumentRequest struct { + ProviderName *string `json:"provider_name"` + InstanceName *string `json:"instance_name"` + ModelName *string `json:"model_name"` + Query string `json:"query"` + Documents []string `json:"documents"` + TopN int `json:"top_n"` +} + +func (h *ProviderHandler) RerankDocument(c *gin.Context) { + var req RerankDocumentRequest + if err := c.ShouldBindJSON(&req); err != nil { + println("JSON bind error: %v (type: %T)", err, err) + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": err.Error(), + }) + return + } + + if req.ProviderName == nil || *req.ProviderName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + if req.InstanceName == nil || *req.InstanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + if req.ModelName == nil || *req.ModelName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model name is required", + }) + return + } + + userID := c.GetString("user_id") + + apiConfig := models.APIConfig{ + ApiKey: nil, + Region: nil, + } + + rerankConfig := models.RerankConfig{ + TopN: req.TopN, + } + + // Non-stream response + var response *models.RerankResponse + var errorCode common.ErrorCode + var err error + + response, errorCode, err = h.modelProviderService.RerankDocument(*req.ProviderName, *req.InstanceName, *req.ModelName, userID, req.Query, req.Documents, &apiConfig, &rerankConfig) + + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "data": response.Data, + "message": "success", + }) +} diff --git a/internal/router/router.go b/internal/router/router.go index 9569277f7d..97c9b90984 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -269,6 +269,8 @@ func (r *Router) Setup(engine *gin.Engine) { provider.POST("/:provider_name/instances/:instance_name/models", r.providerHandler.AddCustomModel) provider.DELETE("/:provider_name/instances/:instance_name/models", r.providerHandler.DropInstanceModels) v1.POST("/chat/completions", r.providerHandler.ChatToModel) + v1.POST("/embeddings", r.providerHandler.EmbedText) + v1.POST("/rerank", r.providerHandler.RerankDocument) } model := v1.Group("/models") diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 953a1b51cf..1a107d4231 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -890,6 +890,238 @@ func (m *ModelProviderService) ChatToModelStreamWithSender(providerName, instanc return common.CodeServerError, errors.New("model is disabled") } +// EmbedText sends texts to the embedding model +func (m *ModelProviderService) EmbedText(providerName, instanceName, modelName, userID string, texts []string, apiConfig *modelModule.APIConfig, modelConfig *modelModule.EmbeddingConfig) (*modelModule.EmbeddingResponse, common.ErrorCode, error) { + if apiConfig == nil { + apiConfig = &modelModule.APIConfig{} + } + if modelConfig == nil { + modelConfig = &modelModule.EmbeddingConfig{} + } + + // Get tenant ID from user + tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") + if err != nil { + return nil, common.CodeServerError, err + } + + if len(tenants) == 0 { + return nil, common.CodeNotFound, errors.New("user has no tenants") + } + + tenantID := tenants[0].TenantID + + // Check if provider exists + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + if err != nil { + return nil, common.CodeServerError, err + } + + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + if err != nil { + return nil, common.CodeServerError, err + } + + modelInfo, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + if err != nil { + providerInfo := dao.GetModelProviderManager().FindProvider(providerName) + if providerInfo == nil { + return nil, common.CodeNotFound, errors.New("provider not found") + } + + var model *entity.Model = nil + model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) + if err != nil { + return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) + } + + if !model.ModelTypeMap["embedding"] { + return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s is not an embedding model", providerName, modelName)) + } + + var extra map[string]string + err = json.Unmarshal([]byte(instance.Extra), &extra) + if err != nil { + return nil, common.CodeServerError, err + } + + region := extra["region"] + apiConfig.Region = ®ion + apiConfig.ApiKey = &instance.APIKey + + var embeddingList [][]float64 + embeddingList, err = providerInfo.ModelDriver.Encode(&modelName, texts, apiConfig, modelConfig) + if err != nil { + return nil, common.CodeServerError, err + } + if embeddingList == nil { + return nil, common.CodeServerError, errors.New("empty embed response") + } + + response := &modelModule.EmbeddingResponse{ + Data: make([]modelModule.EmbeddingResult, len(embeddingList)), + } + for i, embedding := range embeddingList { + response.Data[i] = modelModule.EmbeddingResult{ + Index: i, + Dimension: len(embedding), + //Embedding: embedding, + } + } + + return response, common.CodeSuccess, nil + } + + if modelInfo.Status == "active" { + // For local deployed models + providerInfo := dao.GetModelProviderManager().FindProvider(providerName) + if providerInfo == nil { + return nil, common.CodeNotFound, errors.New("provider not found") + } + + var extra map[string]string + err = json.Unmarshal([]byte(instance.Extra), &extra) + if err != nil { + return nil, common.CodeServerError, err + } + + region := extra["region"] + apiConfig.Region = ®ion + apiConfig.ApiKey = &instance.APIKey + + newURL := map[string]string{ + region: extra["base_url"], + } + newProviderInfo := providerInfo.ModelDriver.NewInstance(newURL) + + var embeddingList [][]float64 + embeddingList, err = newProviderInfo.Encode(&modelName, texts, apiConfig, modelConfig) + if err != nil { + return nil, common.CodeServerError, err + } + if embeddingList == nil { + return nil, common.CodeServerError, errors.New("empty embed response") + } + + response := &modelModule.EmbeddingResponse{ + Data: make([]modelModule.EmbeddingResult, len(embeddingList)), + } + for i, embedding := range embeddingList { + response.Data[i] = modelModule.EmbeddingResult{ + Index: i, + Dimension: len(embedding), + //Embedding: embedding, + } + } + + return response, common.CodeSuccess, nil + } + + return nil, common.CodeServerError, errors.New("model is disabled") +} + +// RerankDocument sends texts to the embedding model +func (m *ModelProviderService) RerankDocument(providerName, instanceName, modelName, userID, query string, documents []string, apiConfig *modelModule.APIConfig, modelConfig *modelModule.RerankConfig) (*modelModule.RerankResponse, common.ErrorCode, error) { + if apiConfig == nil { + apiConfig = &modelModule.APIConfig{} + } + if modelConfig == nil { + modelConfig = &modelModule.RerankConfig{} + } + + // Get tenant ID from user + tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") + if err != nil { + return nil, common.CodeServerError, err + } + + if len(tenants) == 0 { + return nil, common.CodeNotFound, errors.New("user has no tenants") + } + + tenantID := tenants[0].TenantID + + // Check if provider exists + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, providerName) + if err != nil { + return nil, common.CodeServerError, err + } + + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + if err != nil { + return nil, common.CodeServerError, err + } + + modelInfo, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + if err != nil { + providerInfo := dao.GetModelProviderManager().FindProvider(providerName) + if providerInfo == nil { + return nil, common.CodeNotFound, errors.New("provider not found") + } + + var model *entity.Model = nil + model, err = dao.GetModelProviderManager().GetModelByName(providerName, modelName) + if err != nil { + return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s not found", providerName, modelName)) + } + + if !model.ModelTypeMap["rerank"] { + return nil, common.CodeNotFound, errors.New(fmt.Sprintf("provider %s model %s is not an embedding model", providerName, modelName)) + } + + var extra map[string]string + err = json.Unmarshal([]byte(instance.Extra), &extra) + if err != nil { + return nil, common.CodeServerError, err + } + + region := extra["region"] + apiConfig.Region = ®ion + apiConfig.ApiKey = &instance.APIKey + + var response *modelModule.RerankResponse + response, err = providerInfo.ModelDriver.Rerank(&modelName, query, documents, apiConfig, modelConfig) + if err != nil { + return nil, common.CodeServerError, err + } + + return response, common.CodeSuccess, nil + } + + if modelInfo.Status == "active" { + // For local deployed models + providerInfo := dao.GetModelProviderManager().FindProvider(providerName) + if providerInfo == nil { + return nil, common.CodeNotFound, errors.New("provider not found") + } + + var extra map[string]string + err = json.Unmarshal([]byte(instance.Extra), &extra) + if err != nil { + return nil, common.CodeServerError, err + } + + region := extra["region"] + apiConfig.Region = ®ion + apiConfig.ApiKey = &instance.APIKey + + newURL := map[string]string{ + region: extra["base_url"], + } + newProviderInfo := providerInfo.ModelDriver.NewInstance(newURL) + + var response *modelModule.RerankResponse + response, err = newProviderInfo.Rerank(&modelName, query, documents, apiConfig, modelConfig) + if err != nil { + return nil, common.CodeServerError, err + } + + return response, common.CodeSuccess, nil + } + + return nil, common.CodeServerError, errors.New("model is disabled") +} + // GetEmbeddingModel returns an EmbeddingModel wrapper for the given tenant func (m *ModelProviderService) GetEmbeddingModel(tenantID, compositeModelName string) (*modelModule.EmbeddingModel, error) { driver, modelName, apiConfig, maxTokens, err := m.getModelConfig(tenantID, compositeModelName) diff --git a/internal/service/nlp/reranker.go b/internal/service/nlp/reranker.go index f127c10009..2e18d5f89c 100644 --- a/internal/service/nlp/reranker.go +++ b/internal/service/nlp/reranker.go @@ -134,20 +134,20 @@ func RerankByModel( // Calculate token similarity tsim = TokenSimilarity(keywords, insTw, qb) + var modelSim []float64 // Get similarity scores from reranker model - modelSim, err := rerankModel.ModelDriver.Rerank(rerankModel.ModelName, query, docs, rerankModel.APIConfig) + rerankResponse, err := rerankModel.ModelDriver.Rerank(rerankModel.ModelName, query, docs, rerankModel.APIConfig, &models.RerankConfig{}) if err != nil { common.Error("RerankByModel: rerankModel.Rerank failed; falling back to token-only similarity", err) // If model fails, fall back to token similarity only modelSim = make([]float64, len(tsim)) } - if len(modelSim) != chunkCount { - common.Warn("reranker returned mismatched score length; padding/truncating", - zap.Int("got", len(modelSim)), zap.Int("want", chunkCount)) - fixed := make([]float64, chunkCount) - copy(fixed, modelSim) - modelSim = fixed + + loopCount := min(chunkCount, len(rerankResponse.Data)) + for i := 0; i < loopCount; i++ { + modelSim = append(modelSim, rerankResponse.Data[i].RelevanceScore) } + // Combine token similarity with model similarity // Model similarity is treated as vector similarity component sim = make([]float64, chunkCount)