From bb05a8bd7e009276629629633def6443f448381e Mon Sep 17 00:00:00 2001 From: Jin Hai Date: Wed, 29 Apr 2026 17:05:08 +0800 Subject: [PATCH] Update create model instance command (#14441) ### What problem does this PR solve? 1. support command: ``` RAGFlow(user)> create provider 'vllm' instance 'test' key 'test-key' url 'base-url' region 'abc'; SUCCESS RAGFlow(user)> list instances from 'vllm'; +----------+----------------------------------------+----------------------------------+--------------+----------------------------------+--------+ | apiKey | extra | id | instanceName | providerID | status | +----------+----------------------------------------+----------------------------------+--------------+----------------------------------+--------+ | test-key | {"base_url":"base-url","region":"abc"} | 40213c89430311f1a7cf38a74640adcc | test | b4d40e6142d311f1a4f938a74640adcc | enable | +----------+----------------------------------------+----------------------------------+--------------+----------------------------------+--------+ ``` 2. support add vllm model ``` RAGFlow(user)> add model 'Qwen/Qwen2-0.5B' to provider 'vllm' instance 'test' with tokens 131072 chat; SUCCESS ``` 3. add vllm chat ### Type of change - [x] New Feature (non-breaking change which adds functionality) - [x] Refactoring --------- Signed-off-by: Jin Hai --- conf/models/vllm.json | 8 + internal/cli/client.go | 2 + internal/cli/http_client.go | 2 +- internal/cli/lexer.go | 4 + internal/cli/types.go | 2 + internal/cli/user_command.go | 81 +++++++++ internal/cli/user_parser.go | 193 +++++++++++++++++++++- internal/entity/model.go | 9 +- internal/entity/models/aliyun.go | 4 + internal/entity/models/deepseek.go | 4 + internal/entity/models/dummy.go | 4 + internal/entity/models/factory.go | 2 + internal/entity/models/gitee.go | 4 + internal/entity/models/google.go | 4 + internal/entity/models/minimax.go | 4 + internal/entity/models/moonshot.go | 4 + internal/entity/models/siliconflow.go | 4 + internal/entity/models/types.go | 4 +- internal/entity/models/vllm.go | 229 ++++++++++++++++++++++++++ internal/entity/models/volcengine.go | 4 + internal/entity/models/zhipu-ai.go | 4 + internal/handler/providers.go | 66 +++++++- internal/router/router.go | 1 + internal/service/model_service.go | 133 +++++++++++++-- 24 files changed, 754 insertions(+), 22 deletions(-) create mode 100644 conf/models/vllm.json create mode 100644 internal/entity/models/vllm.go diff --git a/conf/models/vllm.json b/conf/models/vllm.json new file mode 100644 index 0000000000..96ec1a2403 --- /dev/null +++ b/conf/models/vllm.json @@ -0,0 +1,8 @@ +{ + "name": "vllm", + "url_suffix": { + "chat": "chat/completions", + "models": "models" + }, + "class": "local" +} \ No newline at end of file diff --git a/internal/cli/client.go b/internal/cli/client.go index 18a0be69ac..acd8eba175 100644 --- a/internal/cli/client.go +++ b/internal/cli/client.go @@ -246,6 +246,8 @@ func (c *RAGFlowClient) ExecuteUserCommand(cmd *Command) (ResponseIf, error) { return c.EnableOrDisableModel(cmd, "enable") case "disable_model": return c.EnableOrDisableModel(cmd, "disable") + case "add_custom_model": + return c.AddCustomModel(cmd) case "chat_to_model": return c.ChatToModel(cmd) case "think_chat_to_model": diff --git a/internal/cli/http_client.go b/internal/cli/http_client.go index cab9858407..6dc1a8846b 100644 --- a/internal/cli/http_client.go +++ b/internal/cli/http_client.go @@ -54,7 +54,7 @@ func NewHTTPClient() *HTTPClient { VerifySSL: false, client: &http.Client{ Transport: transport, - Timeout: 60 * time.Second, + Timeout: 300 * time.Second, }, } } diff --git a/internal/cli/lexer.go b/internal/cli/lexer.go index 4f5c4c1963..c8ffb1bffd 100644 --- a/internal/cli/lexer.go +++ b/internal/cli/lexer.go @@ -415,6 +415,10 @@ func (l *Lexer) lookupIdent(ident string) Token { return Token{Type: TokenDocument, Value: ident} case "TAGS": return Token{Type: TokenTag, Value: ident} + case "REGION": + return Token{Type: TokenRegion, Value: ident} + case "URL": + return Token{Type: TokenURL, Value: ident} case "LOG": return Token{Type: TokenLog, Value: ident} case "LEVEL": diff --git a/internal/cli/types.go b/internal/cli/types.go index 286f310c47..12822f4a64 100644 --- a/internal/cli/types.go +++ b/internal/cli/types.go @@ -137,6 +137,8 @@ const ( TokenChunks TokenDocument TokenTag + TokenRegion + TokenURL TokenLog TokenLevel TokenDebug diff --git a/internal/cli/user_command.go b/internal/cli/user_command.go index 87fca57092..2e30b52adb 100644 --- a/internal/cli/user_command.go +++ b/internal/cli/user_command.go @@ -1129,11 +1129,23 @@ func (c *RAGFlowClient) CreateProviderInstance(cmd *Command) (ResponseIf, error) return nil, fmt.Errorf("API key not provided") } + baseUrl, ok := cmd.Params["base_url"].(string) + if !ok { + baseUrl = "" + } + + region, ok := cmd.Params["region"].(string) + if !ok { + region = "" + } + url := fmt.Sprintf("/providers/%s/instances", providerName) payload := map[string]interface{}{ "instance_name": instanceName, "api_key": apiKey, + "base_url": baseUrl, + "region": region, } resp, err := c.HTTPClient.Request("POST", url, true, "web", nil, payload) @@ -1685,6 +1697,75 @@ func (c *RAGFlowClient) ShowCurrentModel(cmd *Command) (ResponseIf, error) { return &result, nil } +func (c *RAGFlowClient) AddCustomModel(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") + } + + providerName, ok := cmd.Params["provider_name"].(string) + if !ok { + return nil, fmt.Errorf("provider name not provided") + } + + instanceName, ok := cmd.Params["instance_name"].(string) + if !ok { + return nil, fmt.Errorf("instance name not provided") + } + + modelName, ok := cmd.Params["model_name"].(string) + if !ok { + return nil, fmt.Errorf("model name not provided") + } + + // chat, vision, embedding, rerank, tts, asr, ocr + modelType, ok := cmd.Params["model_type"].(string) + if !ok { + return nil, fmt.Errorf("model type not provided") + } + + maxTokens, ok := cmd.Params["max_tokens"].(int) + if !ok { + return nil, fmt.Errorf("max tokens not provided") + } + + url := fmt.Sprintf("/providers/%s/instances/%s/models", providerName, instanceName) + + payload := map[string]interface{}{ + "provider_name": providerName, + "instance_name": instanceName, + "model_name": modelName, + "model_type": modelType, + "max_tokens": maxTokens, + } + + supportThink, ok := cmd.Params["support_think"].(bool) + if ok { + payload["thinking"] = supportThink + } + + resp, err := c.HTTPClient.Request("POST", url, true, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to check provider connection: %w", err) + } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to check provider connection: 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) + } + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + result.Duration = resp.Duration + return &result, nil + +} + // Context related commands // CEList handles the ls command - lists nodes using Context Engine diff --git a/internal/cli/user_parser.go b/internal/cli/user_parser.go index 2db84b55cd..a31a374ec5 100644 --- a/internal/cli/user_parser.go +++ b/internal/cli/user_parser.go @@ -531,6 +531,8 @@ func (p *Parser) parseAddCommand() (*Command, error) { switch p.curToken.Type { case TokenProvider: return p.parseAddProvider() + case TokenModel: + return p.parseAddModel() default: return nil, fmt.Errorf("unknown ADD target: %s", p.curToken.Value) } @@ -721,6 +723,154 @@ func (p *Parser) parseAddProvider() (*Command, error) { return cmd, nil } +// syntax: add model 'xxx' to provider 'vllm' instance 'test' with tokens 1024 chat think vision; +func (p *Parser) parseAddModel() (*Command, error) { + p.nextToken() // consume MODEL + + if p.curToken.Type != TokenQuotedString { + return nil, fmt.Errorf("expected model name") + } + + modelName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() // consume model name + + if p.curToken.Type != TokenTo { + return nil, fmt.Errorf("expected TO") + } + p.nextToken() + + if p.curToken.Type != TokenProvider { + return nil, fmt.Errorf("expected PROVIDER") + } + p.nextToken() + + // provider name + if p.curToken.Type != TokenQuotedString { + return nil, fmt.Errorf("expected provider name") + } + providerName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + if p.curToken.Type != TokenInstance { + return nil, fmt.Errorf("expected INSTANCE") + } + p.nextToken() + + // instance name + if p.curToken.Type != TokenQuotedString { + return nil, fmt.Errorf("expected provider name") + } + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + modelType := "" + var supportThink *bool = nil + maxTokens := 0 + if p.curToken.Type == TokenWith { + p.nextToken() // pass WITH + optionsLoop: + for { + switch p.curToken.Type { + case TokenThink: + if supportThink != nil { + return nil, fmt.Errorf("think model is already set") + } + supportThink = new(bool) + p.nextToken() + *supportThink = true + case TokenVision: + p.nextToken() + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to vision", modelType) + } + modelType = "vision" + case TokenChat: + p.nextToken() + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to chat", modelType) + } + modelType = "chat" + case TokenEmbedding: + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to embedding", modelType) + } + p.nextToken() + modelType = "embedding" + case TokenRerank: + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to rerank", modelType) + } + p.nextToken() + modelType = "rerank" + case TokenOCR: + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to OCR", modelType) + } + p.nextToken() + modelType = "ocr" + case TokenTTS: + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to TTS", modelType) + } + p.nextToken() + modelType = "tts" + case TokenASR: + if modelType != "" { + return nil, fmt.Errorf("model type is %s, attempt to change to ASR", modelType) + } + p.nextToken() + modelType = "asr" + case TokenTokens: + p.nextToken() // pass TOKENS + if maxTokens != 0 { + return nil, fmt.Errorf("max tokens is already given %d", maxTokens) + } + if p.curToken.Type != TokenInteger { + return nil, fmt.Errorf("expected integer") + } + maxTokens, err = p.parseNumber() + if err != nil { + return nil, err + } + p.nextToken() // consume + case TokenSemicolon: + p.nextToken() + break optionsLoop // done + default: + // No more options to process + break optionsLoop + } + } + } + + cmd := NewCommand("add_custom_model") + cmd.Params["model_name"] = modelName + cmd.Params["model_type"] = modelType + cmd.Params["provider_name"] = providerName + cmd.Params["instance_name"] = instanceName + if supportThink != nil { + cmd.Params["support_think"] = *supportThink + } + cmd.Params["max_tokens"] = maxTokens + + if modelType != "chat" && modelType != "vision" { + if supportThink != nil && *supportThink { + return nil, fmt.Errorf("think not supported for model type %s", modelType) + } + } + + return cmd, nil +} + func (p *Parser) parseCreateDataset() (*Command, error) { p.nextToken() // consume DATASET datasetName, err := p.parseQuotedString() @@ -1201,7 +1351,7 @@ func (p *Parser) parseAlterProvider() (*Command, error) { return cmd, nil } -// parseCreateProviderInstance parses CREATE PROVIDER INSTANCE command +// parseCreateProviderInstance parses CREATE PROVIDER INSTANCE KEY URL command // instance_name cannot be "default" func (p *Parser) parseCreateProviderInstance() (*Command, error) { p.nextToken() // consume PROVIDER @@ -1226,17 +1376,54 @@ func (p *Parser) parseCreateProviderInstance() (*Command, error) { if instanceName == "default" { return nil, fmt.Errorf("instance name cannot be 'default'") } - p.nextToken() + + if p.curToken.Type != TokenKey { + return nil, fmt.Errorf("expected KEY after instance name") + } + p.nextToken() + apiKey, err := p.parseQuotedString() if err != nil { return nil, fmt.Errorf("expected API key: %w", err) } + p.nextToken() + + baseURL := "" + if p.curToken.Type == TokenURL { + p.nextToken() + baseURL, err = p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected base URL: %w", err) + } + p.nextToken() + } + + region := "" + if p.curToken.Type == TokenRegion { + p.nextToken() + region, err = p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected base URL: %w", err) + } + p.nextToken() + } cmd := NewCommand("create_provider_instance") cmd.Params["provider_name"] = providerName cmd.Params["instance_name"] = instanceName cmd.Params["api_key"] = apiKey + if baseURL != "" { + // Only local model provider need to set URL + cmd.Params["base_url"] = baseURL + if region == "" { + region = instanceName + } + } + + if region != "" { + cmd.Params["region"] = region + } p.nextToken() // Semicolon is optional @@ -2280,7 +2467,7 @@ func (p *Parser) parseChatCommand() (*Command, error) { switch p.curToken.Type { case TokenEffort: { - p.nextToken() // pass VERBOSITY + p.nextToken() // pass Effort switch p.curToken.Type { case TokenNone: effort = "none" diff --git a/internal/entity/model.go b/internal/entity/model.go index 54a28cc08b..08a2958a5f 100644 --- a/internal/entity/model.go +++ b/internal/entity/model.go @@ -319,22 +319,21 @@ func (pm *ProviderManager) ListModels(providerName string) ([]map[string]interfa return nil, fmt.Errorf("provider '%s' not found", providerName) } - models := []map[string]interface{}{} + modelList := []map[string]interface{}{} for _, model := range provider.Models { modelData := map[string]interface{}{ "name": model.Name, "max_tokens": model.MaxTokens, "model_types": model.ModelTypes, - "features": GetFeatures(model), } - models = append(models, modelData) + modelList = append(modelList, modelData) } - if len(models) == 0 { + if len(modelList) == 0 { return nil, fmt.Errorf("no models found") } - return models, nil + return modelList, nil } func (pm *ProviderManager) GetModelByName(providerName, modelName string) (*Model, error) { diff --git a/internal/entity/models/aliyun.go b/internal/entity/models/aliyun.go index 48ef6b7066..5613e76617 100644 --- a/internal/entity/models/aliyun.go +++ b/internal/entity/models/aliyun.go @@ -52,6 +52,10 @@ func NewAliyunModel(baseURL map[string]string, urlSuffix URLSuffix) *AliyunModel } } +func (z *AliyunModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *AliyunModel) Name() string { return "siliconflow" } diff --git a/internal/entity/models/deepseek.go b/internal/entity/models/deepseek.go index ee47918a54..2e8b894f93 100644 --- a/internal/entity/models/deepseek.go +++ b/internal/entity/models/deepseek.go @@ -52,6 +52,10 @@ func NewDeepSeekModel(baseURL map[string]string, urlSuffix URLSuffix) *DeepSeekM } } +func (z *DeepSeekModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *DeepSeekModel) Name() string { return "deepseek" } diff --git a/internal/entity/models/dummy.go b/internal/entity/models/dummy.go index 59a84b49fe..d02ac04159 100644 --- a/internal/entity/models/dummy.go +++ b/internal/entity/models/dummy.go @@ -34,6 +34,10 @@ func NewDummyModel(baseURL map[string]string, urlSuffix URLSuffix) *DummyModel { } } +func (z *DummyModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *DummyModel) Name() string { return "dummy" } diff --git a/internal/entity/models/factory.go b/internal/entity/models/factory.go index e6e0c5f1da..eb42783fba 100644 --- a/internal/entity/models/factory.go +++ b/internal/entity/models/factory.go @@ -51,6 +51,8 @@ func (f *ModelFactory) CreateModelDriver(providerName string, baseURL map[string return NewAliyunModel(baseURL, urlSuffix), nil case "volcengine": return NewVolcEngine(baseURL, urlSuffix), nil + case "vllm": + return NewVllmModel(baseURL, urlSuffix), nil default: return NewDummyModel(baseURL, urlSuffix), nil } diff --git a/internal/entity/models/gitee.go b/internal/entity/models/gitee.go index b28bedea13..1eca6eb919 100644 --- a/internal/entity/models/gitee.go +++ b/internal/entity/models/gitee.go @@ -52,6 +52,10 @@ func NewGiteeModel(baseURL map[string]string, urlSuffix URLSuffix) *GiteeModel { } } +func (z *GiteeModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *GiteeModel) Name() string { return "gitee" } diff --git a/internal/entity/models/google.go b/internal/entity/models/google.go index cbc42b2812..4adb6490d4 100644 --- a/internal/entity/models/google.go +++ b/internal/entity/models/google.go @@ -38,6 +38,10 @@ func NewGoogleModel(baseURL map[string]string, urlSuffix URLSuffix) *GoogleModel } } +func (z *GoogleModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *GoogleModel) Name() string { return "google" } diff --git a/internal/entity/models/minimax.go b/internal/entity/models/minimax.go index c1001d50c8..9fe32a289d 100644 --- a/internal/entity/models/minimax.go +++ b/internal/entity/models/minimax.go @@ -47,6 +47,10 @@ func NewMinimaxModel(baseURL map[string]string, urlSuffix URLSuffix) *MinimaxMod } } +func (z *MinimaxModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *MinimaxModel) Name() string { return "minimax" } diff --git a/internal/entity/models/moonshot.go b/internal/entity/models/moonshot.go index b436d672f1..a55787f48a 100644 --- a/internal/entity/models/moonshot.go +++ b/internal/entity/models/moonshot.go @@ -52,6 +52,10 @@ func NewMoonshotModel(baseURL map[string]string, urlSuffix URLSuffix) *MoonshotM } } +func (z *MoonshotModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *MoonshotModel) Name() string { return "moonshot" } diff --git a/internal/entity/models/siliconflow.go b/internal/entity/models/siliconflow.go index 2c191b3349..11b59e1d21 100644 --- a/internal/entity/models/siliconflow.go +++ b/internal/entity/models/siliconflow.go @@ -52,6 +52,10 @@ func NewSiliconflowModel(baseURL map[string]string, urlSuffix URLSuffix) *Silico } } +func (z *SiliconflowModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *SiliconflowModel) Name() string { return "siliconflow" } diff --git a/internal/entity/models/types.go b/internal/entity/models/types.go index fd4e031b0a..90a9a69aee 100644 --- a/internal/entity/models/types.go +++ b/internal/entity/models/types.go @@ -8,6 +8,8 @@ type Message struct { // EmbeddingModel interface for embedding models type ModelDriver interface { + NewInstance(baseURL map[string]string) ModelDriver + Name() string // Chat sends a message and returns response @@ -20,7 +22,7 @@ type ModelDriver interface { 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) - // List suppported models + // ListModels List supported models ListModels(apiConfig *APIConfig) ([]string, error) Balance(apiConfig *APIConfig) (map[string]interface{}, error) diff --git a/internal/entity/models/vllm.go b/internal/entity/models/vllm.go new file mode 100644 index 0000000000..6cfdef91b4 --- /dev/null +++ b/internal/entity/models/vllm.go @@ -0,0 +1,229 @@ +// +// 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 models + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "time" +) + +// VllmModel implements ModelDriver for Vllm AI +type VllmModel struct { + BaseURL map[string]string + URLSuffix URLSuffix + httpClient *http.Client // Reusable HTTP client with connection pool +} + +// NewVllmModel creates a new Vllm AI model instance +func NewVllmModel(baseURL map[string]string, urlSuffix URLSuffix) *VllmModel { + return &VllmModel{ + BaseURL: baseURL, + URLSuffix: urlSuffix, + httpClient: &http.Client{ + Timeout: 120 * time.Second, + Transport: &http.Transport{ + MaxIdleConns: 100, + MaxIdleConnsPerHost: 10, + IdleConnTimeout: 90 * time.Second, + DisableCompression: false, + }, + }, + } +} + +func (z *VllmModel) NewInstance(baseURL map[string]string) ModelDriver { + return &VllmModel{ + BaseURL: baseURL, + URLSuffix: z.URLSuffix, + httpClient: &http.Client{ + Timeout: 120 * time.Second, + Transport: &http.Transport{ + MaxIdleConns: 100, + MaxIdleConnsPerHost: 10, + IdleConnTimeout: 90 * time.Second, + DisableCompression: false, + }, + }, + } +} + +func (z *VllmModel) Name() string { + return "vllm" +} + +// Chat sends a message and returns response +func (z *VllmModel) Chat(modelName, message *string, apiConfig *APIConfig, chatModelConfig *ChatConfig) (*ChatResponse, error) { + if message == nil { + return nil, fmt.Errorf("message is nil") + } + + var region = "default" + if apiConfig.Region != nil { + region = *apiConfig.Region + } + + url := fmt.Sprintf("%s/%s", z.BaseURL[region], z.URLSuffix.Chat) + + // I need to get the model type, such as qwen3 is the prefix, the model type will be qwen. glm is the prefix, the model type will be glm. such as the model name: qwen3-0.6b, the model type will be qwen3 + // the model name is glm-4.7, the model type will be glm + modelType := strings.Split(*modelName, "-")[0] + if modelType == "qwen" || modelType == "glm" { + url = fmt.Sprintf("%s/%s", z.BaseURL[region], z.URLSuffix.AsyncChat) + } + + // Build request body + reqBody := map[string]interface{}{ + "model": modelName, + "messages": []map[string]string{ + {"role": "user", "content": *message}, + }, + "stream": false, + "temperature": 1, + } + + if chatModelConfig.Stream != nil { + reqBody["stream"] = *chatModelConfig.Stream + } + + if chatModelConfig.MaxTokens != nil { + reqBody["max_tokens"] = *chatModelConfig.MaxTokens + } + + if chatModelConfig.Temperature != nil { + reqBody["temperature"] = *chatModelConfig.Temperature + } + + if chatModelConfig.TopP != nil { + reqBody["top_p"] = *chatModelConfig.TopP + } + + if chatModelConfig.Stop != nil { + reqBody["stop"] = *chatModelConfig.Stop + } + + if chatModelConfig.Thinking != nil { + if *chatModelConfig.Thinking { + reqBody["thinking"] = map[string]interface{}{ + "type": "enabled", + } + } else { + reqBody["thinking"] = map[string]interface{}{ + "type": "disabled", + } + } + } + + jsonData, err := json.Marshal(reqBody) + if err != nil { + return nil, fmt.Errorf("failed to marshal request: %w", err) + } + + req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) + if err != nil { + return nil, fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiConfig.ApiKey)) + + resp, err := z.httpClient.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("failed to read response: %w", err) + } + + if resp.StatusCode != http.StatusOK { + return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body)) + } + + // Parse response + var result map[string]interface{} + if err = json.Unmarshal(body, &result); err != nil { + return nil, fmt.Errorf("failed to parse response: %w", err) + } + + choices, ok := result["choices"].([]interface{}) + if !ok || len(choices) == 0 { + return nil, fmt.Errorf("no choices in response") + } + + firstChoice, ok := choices[0].(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("invalid choice format") + } + + messageMap, ok := firstChoice["message"].(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("invalid message format") + } + + content, ok := messageMap["content"].(string) + if !ok { + return nil, fmt.Errorf("invalid content format") + } + + thinking, answer := GetThinkingAndAnswer(chatModelConfig.ModelClass, &content) + + chatResponse := &ChatResponse{ + Answer: answer, + ReasonContent: thinking, + } + + return chatResponse, nil +} + +// ChatWithMessages sends multiple messages with roles and returns response +func (z *VllmModel) ChatWithMessages(modelName string, apiKey *string, messages []Message, modelConfig *ChatConfig) (string, error) { + return "", fmt.Errorf("not implemented") +} + +// ChatStreamlyWithSender sends a message and streams response via sender function (best performance, no channel) +func (z *VllmModel) ChatStreamlyWithSender(modelName, message *string, apiConfig *APIConfig, modelConfig *ChatConfig, sender func(*string, *string) error) error { + return fmt.Errorf("not implemented") +} + +// Encode encodes a list of texts into embeddings +func (z *VllmModel) Encode(modelName *string, texts []string, apiConfig *APIConfig, embeddingConfig *EmbeddingConfig) ([][]float64, error) { + return nil, fmt.Errorf("not implemented") +} + +func (z *VllmModel) ListModels(apiConfig *APIConfig) ([]string, error) { + return nil, fmt.Errorf("not implemented") +} + +func (z *VllmModel) Balance(apiConfig *APIConfig) (map[string]interface{}, error) { + return nil, fmt.Errorf("no such method") +} + +func (z *VllmModel) CheckConnection(apiConfig *APIConfig) error { + return fmt.Errorf("no such method") +} + +// Rerank calculates similarity scores between query and texts +func (z *VllmModel) Rerank(modelName *string, query string, texts []string, apiConfig *APIConfig) ([]float64, 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 c189854368..a7fc5b6769 100644 --- a/internal/entity/models/volcengine.go +++ b/internal/entity/models/volcengine.go @@ -52,6 +52,10 @@ func NewVolcEngine(baseURL map[string]string, urlSuffix URLSuffix) *VolcEngine { } } +func (z *VolcEngine) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *VolcEngine) Name() string { return "volcengine" } diff --git a/internal/entity/models/zhipu-ai.go b/internal/entity/models/zhipu-ai.go index cc30578102..ee9ea289ab 100644 --- a/internal/entity/models/zhipu-ai.go +++ b/internal/entity/models/zhipu-ai.go @@ -52,6 +52,10 @@ func NewZhipuAIModel(baseURL map[string]string, urlSuffix URLSuffix) *ZhipuAIMod } } +func (z *ZhipuAIModel) NewInstance(baseURL map[string]string) ModelDriver { + return nil +} + func (z *ZhipuAIModel) Name() string { return "zhipu" } diff --git a/internal/handler/providers.go b/internal/handler/providers.go index 1446a94a82..5104076d77 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -241,7 +241,9 @@ func (h *ProviderHandler) ShowModel(c *gin.Context) { type CreateProviderInstanceRequest struct { InstanceName string `json:"instance_name" binding:"required"` - APIKey string `json:"api_key" binding:"required"` + APIKey string `json:"api_key"` + BaseURL string `json:"base_url"` + Region string `json:"region"` } func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) { @@ -274,7 +276,7 @@ func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) { userID := c.GetString("user_id") - _, err := h.modelProviderService.CreateProviderInstance(providerName, req.InstanceName, req.APIKey, userID, "default") + _, err := h.modelProviderService.CreateProviderInstance(providerName, req.InstanceName, req.APIKey, req.BaseURL, req.Region, userID) if err != nil { c.JSON(http.StatusOK, gin.H{ "code": common.CodeServerError, @@ -645,6 +647,66 @@ func (h *ProviderHandler) EnableOrDisableModel(c *gin.Context) { }) } +func (h *ProviderHandler) AddCustomModel(c *gin.Context) { + var req service.AddCustomModelRequest + 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 == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + if req.InstanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + if req.ModelName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model name is required", + }) + return + } + + if req.ModelType == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model type is required", + }) + return + } + + userID := c.GetString("user_id") + + errorCode, err := h.modelProviderService.AddCustomModel(&req, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeSuccess, + }) + +} + type ChatToModelRequest struct { ProviderName *string `json:"provider_name"` InstanceName *string `json:"instance_name"` diff --git a/internal/router/router.go b/internal/router/router.go index bc33f995c7..ab8c44197e 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -218,6 +218,7 @@ func (r *Router) Setup(engine *gin.Engine) { provider.DELETE("/:provider_name/instances", r.providerHandler.DropProviderInstance) provider.GET("/:provider_name/instances/:instance_name/models", r.providerHandler.ListInstanceModels) provider.PATCH("/:provider_name/instances/:instance_name/models/*model_name", r.providerHandler.EnableOrDisableModel) + provider.POST("/:provider_name/instances/:instance_name/models", r.providerHandler.AddCustomModel) v1.POST("/chat/completions", r.providerHandler.ChatToModel) } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 85edf695bd..043b5ff4d7 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -202,7 +202,7 @@ func (m *ModelProviderService) ListSupportedModels(providerName, instanceName, u return providerInfo.ModelDriver.ListModels(apiConfig) } -func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName, apiKey, userID, region string) (common.ErrorCode, error) { +func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName, apiKey, baseURL, region, userID string) (common.ErrorCode, error) { // Get tenant ID from user tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") if err != nil { @@ -228,6 +228,7 @@ func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName extra := make(map[string]string) extra["region"] = region + extra["base_url"] = baseURL // convert extra to string extraByte, err := json.Marshal(extra) if err != nil { @@ -252,7 +253,7 @@ func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName err = m.modelInstanceDAO.Create(tenantModelProvider) if err != nil { - return common.CodeServerError, errors.New("fail to create model provider") + return common.CodeServerError, fmt.Errorf("fail to create model instance: %s", err.Error()) } return common.CodeSuccess, nil } @@ -298,7 +299,7 @@ func (m *ModelProviderService) ListProviderInstances(providerName, userID string "providerID": instance.ProviderID, "apiKey": instance.APIKey, "status": instance.Status, - "region": extra["region"], + "extra": instance.Extra, }) } @@ -521,23 +522,30 @@ func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, us return nil, err } + allModels, err := dao.GetModelProviderManager().ListModels(providerName) + // insert models name into a set modelNames := make(map[string]bool) for _, model := range disabledModels { - modelNames[model.ModelName] = true - } + if model.Status == "active" { + modelData := map[string]interface{}{ + "name": model.ModelName, + } + allModels = append(allModels, modelData) + } else { + modelNames[model.ModelName] = true + } - allModels, err := dao.GetModelProviderManager().ListModels(providerName) + } for _, model := range allModels { // convert model["name"] to string modelName := model["name"].(string) if modelNames[modelName] { - model["status"] = "disabled" + model["status"] = "inactive" } else { - model["status"] = "enabled" + model["status"] = "active" } - } return allModels, nil @@ -634,7 +642,7 @@ func (m *ModelProviderService) ChatToModel(providerName, instanceName, modelName return nil, common.CodeServerError, err } - _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + modelInfo, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) if err != nil { providerInfo := dao.GetModelProviderManager().FindProvider(providerName) if providerInfo == nil { @@ -668,6 +676,38 @@ func (m *ModelProviderService) ChatToModel(providerName, instanceName, modelName 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 + + modelConfig.ModelClass = &providerInfo.Class + + newURL := map[string]string{ + region: extra["base_url"], + } + newProviderInfo := providerInfo.ModelDriver.NewInstance(newURL) + + var response *modelModule.ChatResponse + response, err = newProviderInfo.Chat(&modelName, &message, apiConfig, modelConfig) + if err != nil { + return nil, common.CodeServerError, err + } + return response, common.CodeSuccess, nil + } + return nil, common.CodeServerError, errors.New("model is disabled") } @@ -850,6 +890,79 @@ func (m *ModelProviderService) GetChatModel(tenantID, compositeModelName string) return modelModule.NewChatModel(driver, &modelName, apiConfig), nil } +type AddCustomModelRequest struct { + ProviderName string `json:"provider_name"` + InstanceName string `json:"instance_name"` + ModelName string `json:"model_name"` + ModelType string `json:"model_type"` + MaxTokens int `json:"max_tokens"` + Thinking *bool `json:"thinking"` +} + +func (m *ModelProviderService) AddCustomModel(request *AddCustomModelRequest, userID string) (common.ErrorCode, error) { + // Get tenant ID from user + tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") + if err != nil { + return common.CodeServerError, err + } + + if len(tenants) == 0 { + return common.CodeNotFound, errors.New("user has no tenants") + } + + tenantID := tenants[0].TenantID + + // Check if provider exists + provider, err := m.modelProviderDAO.GetByTenantIDAndProviderName(tenantID, request.ProviderName) + if err != nil { + return common.CodeServerError, err + } + + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, request.InstanceName) + if err != nil { + return common.CodeServerError, err + } + + _, err = m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, request.ModelName) + if err == nil { + return common.CodeConflict, errors.New("model already exists") + } + + modelID, err := generateUUID1Hex() + if err != nil { + return common.CodeServerError, errors.New("fail to get UUID") + } + + extra := make(map[string]interface{}) + extra["max_tokens"] = request.MaxTokens + if request.Thinking != nil { + extra["thinking"] = *request.Thinking + } + // convert extra to string + extraByte, err := json.Marshal(extra) + if err != nil { + return common.CodeServerError, errors.New("fail to marshal extra") + } + extraStr := string(extraByte) + + model := &entity.TenantModel{ + ID: modelID, + ModelName: request.ModelName, + ModelType: request.ModelType, + ProviderID: provider.ID, + InstanceID: instance.ID, + Status: "active", + Extra: extraStr, + } + + err = m.modelDAO.Create(model) + if err != nil { + return common.CodeServerError, err + } + + return common.CodeSuccess, nil +} + // getModelConfig returns the model driver, model name, and API config for a model func (m *ModelProviderService) getModelConfig(tenantID, compositeModelName string) (modelModule.ModelDriver, string, *modelModule.APIConfig, error) { modelName, instanceName, providerName, err := parseModelName(compositeModelName)