From 6c29128de13083fb3a9f3ff1587fff8e2a806784 Mon Sep 17 00:00:00 2001 From: Jin Hai Date: Thu, 2 Apr 2026 20:20:35 +0800 Subject: [PATCH] Refactor model provider and command (#13887) ### What problem does this PR solve? Introduce 5 new tables, including model groups and provider instance. ### Type of change - [x] New Feature (non-breaking change which adds functionality) - [x] Refactoring --------- Signed-off-by: Jin Hai --- cmd/server_main.go | 3 +- conf/models/openai.json | 3 + conf/models/xai.json | 3 + conf/models/zhipu-ai.json | 182 +++++++ internal/cli/admin_parser.go | 24 - internal/cli/cli.go | 12 +- internal/cli/client.go | 43 +- internal/cli/lexer.go | 16 +- internal/cli/parser.go | 12 + internal/cli/response.go | 28 ++ internal/cli/types.go | 7 + internal/cli/user_command.go | 448 ++++++++++++++++- internal/cli/user_parser.go | 467 +++++++++++++++++- internal/dao/database.go | 5 + internal/dao/tenant_model.go | 67 +++ internal/dao/tenant_model_group.go | 39 ++ internal/dao/tenant_model_group_mapping.go | 39 ++ internal/dao/tenant_model_instance.go | 66 +++ internal/dao/tenant_model_provider.go | 74 +++ internal/entity/model.go | 87 +++- internal/entity/models/dummy.go | 50 ++ internal/entity/models/factory.go | 41 ++ internal/entity/models/types.go | 20 + internal/entity/models/zhipu-ai.go | 307 ++++++++++++ internal/entity/tenant_model.go | 33 ++ internal/entity/tenant_model_group.go | 31 ++ internal/entity/tenant_model_group_mapping.go | 33 ++ internal/entity/tenant_model_instance.go | 32 ++ internal/entity/tenant_model_provider.go | 30 ++ internal/handler/providers.go | 445 ++++++++++++++++- internal/router/router.go | 10 + internal/service/kb.go | 2 +- internal/service/model_service.go | 427 ++++++++++++++++ 33 files changed, 2986 insertions(+), 100 deletions(-) create mode 100644 conf/models/zhipu-ai.json create mode 100644 internal/dao/tenant_model.go create mode 100644 internal/dao/tenant_model_group.go create mode 100644 internal/dao/tenant_model_group_mapping.go create mode 100644 internal/dao/tenant_model_instance.go create mode 100644 internal/dao/tenant_model_provider.go create mode 100644 internal/entity/models/dummy.go create mode 100644 internal/entity/models/factory.go create mode 100644 internal/entity/models/types.go create mode 100644 internal/entity/models/zhipu-ai.go create mode 100644 internal/entity/tenant_model.go create mode 100644 internal/entity/tenant_model_group.go create mode 100644 internal/entity/tenant_model_group_mapping.go create mode 100644 internal/entity/tenant_model_instance.go create mode 100644 internal/entity/tenant_model_provider.go diff --git a/cmd/server_main.go b/cmd/server_main.go index 8b90664bb7..d5897febbb 100644 --- a/cmd/server_main.go +++ b/cmd/server_main.go @@ -176,6 +176,7 @@ func startServer(config *server.Config) { searchService := service.NewSearchService() fileService := service.NewFileService() memoryService := service.NewMemoryService() + modelProviderService := service.NewModelProviderService() // Initialize handler layer authHandler := handler.NewAuthHandler() @@ -193,7 +194,7 @@ func startServer(config *server.Config) { searchHandler := handler.NewSearchHandler(searchService, userService) fileHandler := handler.NewFileHandler(fileService, userService) memoryHandler := handler.NewMemoryHandler(memoryService) - providerHandler := handler.NewProviderHandler(userService) + providerHandler := handler.NewProviderHandler(userService, modelProviderService) // Initialize router r := router.NewRouter(authHandler, userHandler, tenantHandler, documentHandler, datasetsHandler, systemHandler, kbHandler, chunkHandler, llmHandler, chatHandler, chatSessionHandler, connectorHandler, searchHandler, fileHandler, memoryHandler, providerHandler) diff --git a/conf/models/openai.json b/conf/models/openai.json index 57bbccb41f..e7c8a61d40 100644 --- a/conf/models/openai.json +++ b/conf/models/openai.json @@ -2,6 +2,9 @@ "name": "OpenAI", "tags": "LLM,TEXT EMBEDDING,TTS,TEXT RE-RANK,SPEECH2TEXT,MODERATION", "url": "https://api.openai.com/v1", + "url_suffix": { + "chat": "chat/completions" + }, "models": [ { "name": "gpt-5.2-pro", diff --git a/conf/models/xai.json b/conf/models/xai.json index af6905ed9f..455069140b 100644 --- a/conf/models/xai.json +++ b/conf/models/xai.json @@ -2,6 +2,9 @@ "name": "xAI", "tags": "LLM", "url": "https://api.x.ai/v1", + "url_suffix": { + "chat": "chat/completions" + }, "models": [ { "name": "grok-4", diff --git a/conf/models/zhipu-ai.json b/conf/models/zhipu-ai.json new file mode 100644 index 0000000000..34f856c139 --- /dev/null +++ b/conf/models/zhipu-ai.json @@ -0,0 +1,182 @@ +{ + "name": "ZHIPU-AI", + "tags": "LLM,TEXT EMBEDDING,SPEECH2TEXT,MODERATION", + "url": "https://open.bigmodel.cn/api/paas/v4", + "url_suffix": { + "chat": "chat/completions", + "async_chat": "async/chat/completions", + "async_result": "async-result", + "embedding": "embedding", + "rerank": "rerank" + }, + "models": [ + { + "name": "glm-4.7", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5-x", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5-air", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5-airx", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5-flash", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4.5v", + "max_tokens": 64000, + "model_types": [ + "image2text" + ], + "features": {} + }, + { + "name": "glm-4-plus", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-0520", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-airx", + "max_tokens": 8000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-air", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-flash", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-flashx", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4-long", + "max_tokens": 1000000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-3-turbo", + "max_tokens": 128000, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "glm-4v", + "max_tokens": 2000, + "model_types": [ + "image2text" + ], + "features": {} + }, + { + "name": "glm-4-9b", + "max_tokens": 8192, + "model_types": [ + "chat" + ], + "features": {} + }, + { + "name": "embedding-2", + "max_tokens": 512, + "model_types": [ + "embedding" + ], + "features": {} + }, + { + "name": "embedding-3", + "max_tokens": 512, + "model_types": [ + "embedding" + ], + "features": {} + }, + { + "name": "glm-asr", + "max_tokens": 4096, + "model_types": [ + "speech2text" + ], + "features": {} + } + ] +} \ No newline at end of file diff --git a/internal/cli/admin_parser.go b/internal/cli/admin_parser.go index 17ecd68cc5..d2b3ac5654 100644 --- a/internal/cli/admin_parser.go +++ b/internal/cli/admin_parser.go @@ -275,30 +275,6 @@ func (p *Parser) parseAdminListDefaultModels() (*Command, error) { return NewCommand("list_user_default_models"), nil } -func (p *Parser) parseListModelsOfProvider() (*Command, error) { - if p.curToken.Type != TokenModels { - return nil, fmt.Errorf("expected MODELS") - } - - p.nextToken() - if p.curToken.Type != TokenFrom { - return nil, fmt.Errorf("expected FROM") - } - p.nextToken() - providerName, err := p.parseQuotedString() - if err != nil { - return nil, err - } - cmd := NewCommand("list_provider_models") - cmd.Params["provider_name"] = providerName - p.nextToken() - // Semicolon is optional for UNSET TOKEN - if p.curToken.Type == TokenSemicolon { - p.nextToken() - } - return cmd, nil -} - func (p *Parser) parseCommonListProviders() (*Command, error) { p.nextToken() // consume AVAILABLE diff --git a/internal/cli/cli.go b/internal/cli/cli.go index fb48bb3e59..198cb92ff1 100644 --- a/internal/cli/cli.go +++ b/internal/cli/cli.go @@ -332,7 +332,7 @@ func looksLikeSQL(s string) bool { "LIST ", "SHOW ", "CREATE ", "DROP ", "ALTER ", "LOGIN ", "REGISTER ", "PING", "GRANT ", "REVOKE ", "SET ", "UNSET ", "UPDATE ", "DELETE ", "INSERT ", - "SELECT ", "DESCRIBE ", "EXPLAIN ", + "SELECT ", "DESCRIBE ", "EXPLAIN ", "ADD ", "ENABLE ", "DISABLE ", "CHAT ", "USE", } for _, prefix := range sqlPrefixes { if strings.HasPrefix(s, prefix) { @@ -1008,15 +1008,19 @@ Commands (User Mode): LIST TOKENS; - List API tokens LIST PROVIDERS; - List available LLM providers CREATE TOKEN; - Create new API token - CREATE PROVIDER 'name'; - Create a provider without API key - CREATE PROVIDER 'name' 'api_key'; - Create a provider with API key + ADD PROVIDER 'name'; - Create a provider without API key + ADD PROVIDER 'name' 'api_key'; - Create a provider with API key DROP TOKEN 'token_value'; - Delete an API token - DROP PROVIDER 'name'; - Delete a provider + DELETE PROVIDER 'name'; - Delete a provider SET TOKEN 'token_value'; - Set and validate API token SHOW TOKEN; - Show current API token SHOW PROVIDER 'name'; - Show provider details + SHOW CURRENT MODEL; - Show current model settings UNSET TOKEN; - Remove current API token ALTER PROVIDER 'name' NAME 'new_name'; - Rename a provider + USE MODEL 'provider/instance/model'; - Set current model for chat + CHAT 'message'; - Chat using current model + CHAT 'provider/instance/model' 'message'; - Chat with specified model Context Engine Commands (no quotes): ls [path] - List resources diff --git a/internal/cli/client.go b/internal/cli/client.go index e609fb8000..3054db7d66 100644 --- a/internal/cli/client.go +++ b/internal/cli/client.go @@ -24,6 +24,13 @@ import ( // PasswordPromptFunc is a function type for password input type PasswordPromptFunc func(prompt string) (string, error) +// CurrentModel holds the current model configuration +type CurrentModel struct { + Provider string + Instance string + Model string +} + // RAGFlowClient handles API interactions with the RAGFlow server type RAGFlowClient struct { HTTPClient *HTTPClient @@ -31,6 +38,7 @@ type RAGFlowClient struct { PasswordPrompt PasswordPromptFunc // Function for password input OutputFormat OutputFormat // Output format: table, plain, json ContextEngine *ce.Engine // Context Engine for virtual filesystem + CurrentModel *CurrentModel // Current model configuration } // NewRAGFlowClient creates a new RAGFlow client @@ -158,6 +166,8 @@ func (c *RAGFlowClient) ExecuteAdminCommand(cmd *Command) (ResponseIf, error) { return c.ShowProvider(cmd) case "list_provider_models": return c.ListModels(cmd) + case "list_instance_models": + return c.ListInstanceModels(cmd) case "show_model": return c.ShowModel(cmd) // TODO: Implement other commands @@ -203,21 +213,44 @@ func (c *RAGFlowClient) ExecuteUserCommand(cmd *Command) (ResponseIf, error) { return c.CreateDocMetaIndex(cmd) case "drop_doc_meta_index": return c.DropDocMetaIndex(cmd) - case "list_pool_providers": + case "list_available_providers": return c.ListAvailableProviders(cmd) case "show_provider": return c.ShowProvider(cmd) case "list_provider_models": return c.ListModels(cmd) + case "list_instance_models": + return c.ListInstanceModels(cmd) case "show_model": return c.ShowModel(cmd) // Provider commands - case "create_provider": - return c.CreateProvider(cmd) + case "add_provider": + return c.AddProvider(cmd) case "list_providers": return c.ListProviders(cmd) - case "drop_provider": - return c.DropProvider(cmd) + case "delete_provider": + return c.DeleteProvider(cmd) + // Provider instance commands + case "create_provider_instance": + return c.CreateProviderInstance(cmd) + case "list_provider_instances": + return c.ListProviderInstances(cmd) + case "show_provider_instance": + return c.ShowProviderInstance(cmd) + case "alter_provider_instance": + return c.AlterProviderInstance(cmd) + case "drop_provider_instance": + return c.DropProviderInstance(cmd) + case "enable_model": + return c.EnableOrDisableModel(cmd, "enable") + case "disable_model": + return c.EnableOrDisableModel(cmd, "disable") + case "chat_to_model": + return c.ChatToModel(cmd) + case "use_model": + return c.UseModel(cmd) + case "show_current_model": + return c.ShowCurrentModel(cmd) // ContextEngine commands case "ce_ls": return c.CEList(cmd) diff --git a/internal/cli/lexer.go b/internal/cli/lexer.go index bfb97de9e6..cc8d6c6d4a 100644 --- a/internal/cli/lexer.go +++ b/internal/cli/lexer.go @@ -181,6 +181,10 @@ func (l *Lexer) lookupIdent(ident string) Token { return Token{Type: TokenActive, Value: ident} case "ADMIN": return Token{Type: TokenAdmin, Value: ident} + case "ADD": + return Token{Type: TokenAdd, Value: ident} + case "DELETE": + return Token{Type: TokenDelete, Value: ident} case "PASSWORD": return Token{Type: TokenPassword, Value: ident} case "DATASET": @@ -305,14 +309,22 @@ func (l *Lexer) lookupIdent(ident string) Token { return Token{Type: TokenAvailable, Value: ident} case "NAME": return Token{Type: TokenName, Value: ident} - case "POOL": - return Token{Type: TokenPool, Value: ident} + case "INSTANCE": + return Token{Type: TokenInstance, Value: ident} + case "INSTANCES": + return Token{Type: TokenInstances, Value: ident} + case "DISABLE": + return Token{Type: TokenDisable, Value: ident} + case "ENABLE": + return Token{Type: TokenEnable, Value: ident} case "INSERT": return Token{Type: TokenInsert, Value: ident} case "FILE": return Token{Type: TokenFile, Value: ident} case "METADATA": return Token{Type: TokenMetadata, Value: ident} + case "USE": + return Token{Type: TokenUse, Value: ident} default: return Token{Type: TokenIdentifier, Value: ident} } diff --git a/internal/cli/parser.go b/internal/cli/parser.go index 98088e782f..cb26220252 100644 --- a/internal/cli/parser.go +++ b/internal/cli/parser.go @@ -150,6 +150,10 @@ func (p *Parser) parseUserCommand() (*Command, error) { return p.parseCreateCommand() case TokenDrop: return p.parseDropCommand() + case TokenAdd: + return p.parseAddCommand() + case TokenDelete: + return p.parseDeleteCommand() case TokenAlter: return p.parseAlterCommand() case TokenGrant: @@ -182,6 +186,14 @@ func (p *Parser) parseUserCommand() (*Command, error) { return p.parseShutdownCommand() case TokenRestart: return p.parseRestartCommand() + case TokenEnable: + return p.parseEnableCommand() + case TokenDisable: + return p.parseDisableCommand() + case TokenChat: + return p.parseChatCommand() + case TokenUse: + return p.parseUseCommand() default: return nil, fmt.Errorf("unknown command: %s", p.curToken.Value) } diff --git a/internal/cli/response.go b/internal/cli/response.go index 54c565feed..16934aa0e4 100644 --- a/internal/cli/response.go +++ b/internal/cli/response.go @@ -113,6 +113,34 @@ func (r *SimpleResponse) PrintOut() { } } +type MessageResponse struct { + Code int `json:"code"` + Message string `json:"message"` + Duration float64 + outputFormat OutputFormat +} + +func (r *MessageResponse) Type() string { + return "message" +} + +func (r *MessageResponse) TimeCost() float64 { + return r.Duration +} + +func (r *MessageResponse) SetOutputFormat(format OutputFormat) { + r.outputFormat = format +} + +func (r *MessageResponse) PrintOut() { + if r.Code == 0 { + fmt.Println(r.Message) + } else { + fmt.Println("ERROR") + fmt.Printf("%d, %s\n", r.Code, r.Message) + } +} + type RegisterResponse struct { Code int `json:"code"` Message string `json:"message"` diff --git a/internal/cli/types.go b/internal/cli/types.go index b3bd2e6c12..d1f5826056 100644 --- a/internal/cli/types.go +++ b/internal/cli/types.go @@ -42,6 +42,8 @@ const ( TokenAlter TokenActive TokenAdmin + TokenAdd + TokenDelete TokenPassword TokenDataset TokenDatasets @@ -104,6 +106,11 @@ const ( TokenVectorSize TokenDocMeta TokenName // For ALTER PROVIDER NAME + TokenInstance + TokenInstances + TokenDisable + TokenEnable + TokenUse TokenInsert TokenFile TokenMetadata diff --git a/internal/cli/user_command.go b/internal/cli/user_command.go index b5bef94b3c..90dd160d2e 100644 --- a/internal/cli/user_command.go +++ b/internal/cli/user_command.go @@ -751,10 +751,10 @@ func (c *RAGFlowClient) DropDocMetaIndex(cmd *Command) (ResponseIf, error) { return &result, nil } -// CreateProvider creates a new model provider -// CREATE PROVIDER -// CREATE PROVIDER -func (c *RAGFlowClient) CreateProvider(cmd *Command) (ResponseIf, error) { +// AddProvider creates a new model provider +// ADD PROVIDER +// ADD PROVIDER +func (c *RAGFlowClient) AddProvider(cmd *Command) (ResponseIf, error) { if c.ServerType != "user" { return nil, fmt.Errorf("this command is only allowed in USER mode") } @@ -764,28 +764,23 @@ func (c *RAGFlowClient) CreateProvider(cmd *Command) (ResponseIf, error) { return nil, fmt.Errorf("provider name not provided") } - // Get optional api_key - apiKey, _ := cmd.Params["api_key"].(string) - // Build payload payload := map[string]interface{}{ - "llm_factory": providerName, - "api_key": apiKey, - "verify": apiKey != "", // Only verify if api_key is provided + "provider_name": providerName, } - resp, err := c.HTTPClient.Request("POST", "/llm/set_api_key", true, "web", nil, payload) + resp, err := c.HTTPClient.Request("POST", "/providers", true, "web", nil, payload) if err != nil { - return nil, fmt.Errorf("failed to create provider: %w", err) + return nil, fmt.Errorf("failed to add provider: %w", err) } if resp.StatusCode != 200 { - return nil, fmt.Errorf("failed to create provider: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + return nil, fmt.Errorf("failed to add provider: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) } - var result CommonDataResponse + var result SimpleResponse if err = json.Unmarshal(resp.Body, &result); err != nil { - return nil, fmt.Errorf("create provider failed: invalid JSON (%w)", err) + return nil, fmt.Errorf("add provider failed: invalid JSON (%w)", err) } if result.Code != 0 { @@ -803,7 +798,7 @@ func (c *RAGFlowClient) ListProviders(cmd *Command) (ResponseIf, error) { return nil, fmt.Errorf("this command is only allowed in USER mode") } - resp, err := c.HTTPClient.Request("GET", "/llm/factories", true, "web", nil, nil) + resp, err := c.HTTPClient.Request("GET", "/providers", true, "web", nil, nil) if err != nil { return nil, fmt.Errorf("failed to list providers: %w", err) } @@ -825,9 +820,9 @@ func (c *RAGFlowClient) ListProviders(cmd *Command) (ResponseIf, error) { return &result, nil } -// DropProvider deletes a provider -// DROP PROVIDER -func (c *RAGFlowClient) DropProvider(cmd *Command) (ResponseIf, error) { +// DeleteProvider deletes a provider +// DELETE PROVIDER +func (c *RAGFlowClient) DeleteProvider(cmd *Command) (ResponseIf, error) { if c.ServerType != "user" { return nil, fmt.Errorf("this command is only allowed in USER mode") } @@ -837,23 +832,25 @@ func (c *RAGFlowClient) DropProvider(cmd *Command) (ResponseIf, error) { return nil, fmt.Errorf("provider name not provided") } + url := fmt.Sprintf("/providers/%s", providerName) + // Build payload payload := map[string]interface{}{ "llm_factory": providerName, } - resp, err := c.HTTPClient.Request("DELETE", "/llm/factory", true, "web", nil, payload) + resp, err := c.HTTPClient.Request("DELETE", url, true, "web", nil, payload) if err != nil { - return nil, fmt.Errorf("failed to drop provider: %w", err) + return nil, fmt.Errorf("failed to delete provider: %w", err) } if resp.StatusCode != 200 { - return nil, fmt.Errorf("failed to drop provider: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + return nil, fmt.Errorf("failed to delete provider: 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("drop provider failed: invalid JSON (%w)", err) + return nil, fmt.Errorf("delete provider failed: invalid JSON (%w)", err) } if result.Code != 0 { @@ -864,6 +861,410 @@ func (c *RAGFlowClient) DropProvider(cmd *Command) (ResponseIf, error) { return &result, nil } +// CreateProviderInstance creates a new provider instance +// CREATE PROVIDER INSTANCE +func (c *RAGFlowClient) CreateProviderInstance(cmd *Command) (ResponseIf, error) { + 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") + } + + apiKey, ok := cmd.Params["api_key"].(string) + if !ok { + return nil, fmt.Errorf("API key not provided") + } + + url := fmt.Sprintf("/providers/%s/instances", providerName) + + payload := map[string]interface{}{ + "instance_name": instanceName, + "api_key": apiKey, + } + + resp, err := c.HTTPClient.Request("POST", url, true, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to create provider instance: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to create provider instance: 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("create provider instance failed: invalid JSON (%w)", err) + } + + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + + result.Duration = resp.Duration + return &result, nil +} + +// ListProviderInstances lists all instances of a provider +// LIST INSTANCES FROM PROVIDER +func (c *RAGFlowClient) ListProviderInstances(cmd *Command) (ResponseIf, error) { + 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") + } + + url := fmt.Sprintf("/providers/%s/instances", providerName) + + resp, err := c.HTTPClient.Request("GET", url, true, "web", nil, nil) + if err != nil { + return nil, fmt.Errorf("failed to list instances: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to list instances: 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("list instances failed: invalid JSON (%w)", err) + } + + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + + result.Duration = resp.Duration + return &result, nil +} + +// ShowProviderInstance shows details of a specific instance +// SHOW INSTANCE FROM PROVIDER +func (c *RAGFlowClient) ShowProviderInstance(cmd *Command) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + instanceName, ok := cmd.Params["instance_name"].(string) + if !ok { + return nil, fmt.Errorf("instance name not provided") + } + + providerName, ok := cmd.Params["provider_name"].(string) + if !ok { + return nil, fmt.Errorf("provider name not provided") + } + + url := fmt.Sprintf("/providers/%s/instances/%s", providerName, instanceName) + + resp, err := c.HTTPClient.Request("GET", url, true, "web", nil, nil) + if err != nil { + return nil, fmt.Errorf("failed to show instance: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to show instance: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + } + + var result CommonDataResponse + if err = json.Unmarshal(resp.Body, &result); err != nil { + return nil, fmt.Errorf("show instance failed: invalid JSON (%w)", err) + } + + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + + result.Duration = resp.Duration + return &result, nil +} + +// AlterProviderInstance renames a provider instance +// ALTER INSTANCE NAME FROM PROVIDER +func (c *RAGFlowClient) AlterProviderInstance(cmd *Command) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + instanceName, ok := cmd.Params["instance_name"].(string) + if !ok { + return nil, fmt.Errorf("instance name not provided") + } + + newName, ok := cmd.Params["new_name"].(string) + if !ok { + return nil, fmt.Errorf("new name not provided") + } + + providerName, ok := cmd.Params["provider_name"].(string) + if !ok { + return nil, fmt.Errorf("provider name not provided") + } + + url := fmt.Sprintf("/providers/%s/instances/%s", providerName, instanceName) + + payload := map[string]interface{}{ + "llm_name": newName, + } + + resp, err := c.HTTPClient.Request("PUT", url, true, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to alter instance: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to alter instance: 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("alter instance failed: invalid JSON (%w)", err) + } + + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + + result.Duration = resp.Duration + return &result, nil +} + +// DropProviderInstance deletes a provider instance +// DROP INSTANCE FROM PROVIDER +func (c *RAGFlowClient) DropProviderInstance(cmd *Command) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + instanceName, ok := cmd.Params["instance_name"].(string) + if !ok { + return nil, fmt.Errorf("instance name not provided") + } + + providerName, ok := cmd.Params["provider_name"].(string) + if !ok { + return nil, fmt.Errorf("provider name not provided") + } + + url := fmt.Sprintf("/providers/%s/instances/%s", providerName, instanceName) + + resp, err := c.HTTPClient.Request("DELETE", url, true, "web", nil, nil) + if err != nil { + return nil, fmt.Errorf("failed to drop instance: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to drop instance: 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("drop instance 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) ListInstanceModels(cmd *Command) (ResponseIf, error) { + 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") + } + + var endPoint string + endPoint = fmt.Sprintf("/providers/%s/instances/%s/models", providerName, instanceName) + + resp, err := c.HTTPClient.Request("GET", endPoint, true, "web", nil, nil) + if err != nil { + return nil, fmt.Errorf("failed to list instance models: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to list instance models: 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("failed to list instance models: 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) EnableOrDisableModel(cmd *Command, status string) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + modelName, ok := cmd.Params["model_name"].(string) + if !ok { + return nil, fmt.Errorf("model name not provided") + } + + instanceName, ok := cmd.Params["instance_name"].(string) + if !ok { + return nil, fmt.Errorf("instance name not provided") + } + + providerName, ok := cmd.Params["provider_name"].(string) + if !ok { + return nil, fmt.Errorf("provider name not provided") + } + + url := fmt.Sprintf("/providers/%s/instances/%s/models/%s", providerName, instanceName, modelName) + + payload := map[string]interface{}{ + "status": status, + } + + resp, err := c.HTTPClient.Request("PUT", url, true, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to enable/disable model: %w", err) + } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to enable/disable 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("enable/disable model 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) ChatToModel(cmd *Command) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + var providerName, instanceName, modelName string + + // Check if model_name is provided in command + if compositeModelName, ok := cmd.Params["model_name"].(string); ok && compositeModelName != "" { + names := strings.Split(compositeModelName, "/") + if len(names) != 3 { + return nil, fmt.Errorf("model name must be in format 'provider/instance/model'") + } + providerName = names[0] + instanceName = names[1] + modelName = names[2] + } 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") + } + + message := cmd.Params["message"].(string) + + url := fmt.Sprintf("/providers/%s/instances/%s/models/%s", providerName, instanceName, modelName) + + payload := map[string]interface{}{ + "message": message, + } + + resp, err := c.HTTPClient.Request("POST", url, true, "web", nil, payload) + if err != nil { + return nil, fmt.Errorf("failed to chat model: %w", err) + } + if resp.StatusCode != 200 { + return nil, fmt.Errorf("failed to chat model: HTTP %d, body: %s", resp.StatusCode, string(resp.Body)) + } + var result MessageResponse + if err = json.Unmarshal(resp.Body, &result); err != nil { + return nil, fmt.Errorf("chat model failed: invalid JSON (%w)", err) + } + if result.Code != 0 { + return nil, fmt.Errorf("%s", result.Message) + } + result.Duration = resp.Duration + return &result, nil +} + +// UseModel sets the current model for chat +func (c *RAGFlowClient) UseModel(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") + } + + modelIdentifier, ok := cmd.Params["model_identifier"].(string) + if !ok || modelIdentifier == "" { + return nil, fmt.Errorf("model identifier not provided") + } + + names := strings.Split(modelIdentifier, "/") + if len(names) != 3 { + return nil, fmt.Errorf("model identifier must be in format 'provider/instance/model'") + } + + c.CurrentModel = &CurrentModel{ + Provider: names[0], + Instance: names[1], + Model: names[2], + } + + var result SimpleResponse + result.Code = 0 + result.Message = fmt.Sprintf("Current model set to: %s/%s/%s", c.CurrentModel.Provider, c.CurrentModel.Instance, c.CurrentModel.Model) + return &result, nil +} + +// ShowCurrentModel displays the current model configuration +func (c *RAGFlowClient) ShowCurrentModel(cmd *Command) (ResponseIf, error) { + if c.ServerType != "user" { + return nil, fmt.Errorf("this command is only allowed in USER mode") + } + + if c.CurrentModel == nil { + return nil, fmt.Errorf("no current model set. Use 'use model' command first") + } + + var result CommonResponse + result.Code = 0 + result.Data = []map[string]interface{}{ + { + "provider": c.CurrentModel.Provider, + "instance": c.CurrentModel.Instance, + "model": c.CurrentModel.Model, + }, + } + return &result, nil +} + // Context related commands // CEList handles the ls command - lists nodes using Context Engine @@ -947,7 +1348,6 @@ func (c *RAGFlowClient) InsertDatasetFromFile(cmd *Command) (ResponseIf, error) if c.ServerType != "user" { return nil, fmt.Errorf("this command is only allowed in USER mode") } - filePath, ok := cmd.Params["file_path"].(string) if !ok { return nil, fmt.Errorf("file_path not provided") diff --git a/internal/cli/user_parser.go b/internal/cli/user_parser.go index 6411c8e5f7..2069fa3fde 100644 --- a/internal/cli/user_parser.go +++ b/internal/cli/user_parser.go @@ -3,6 +3,7 @@ package cli import ( "fmt" "strconv" + "strings" ) // Command parsers @@ -170,6 +171,8 @@ func (p *Parser) parseListCommand() (*Command, error) { return p.parseListModelsOfProvider() case TokenProviders: return p.parseListProviders() + case TokenInstances: + return p.parseListInstances() case TokenDefault: return p.parseListDefaultModels() case TokenAvailable: @@ -348,15 +351,23 @@ func (p *Parser) parseShowCommand() (*Command, error) { return NewCommand("show_token"), nil case TokenCurrent: p.nextToken() - if p.curToken.Type != TokenUser { - return nil, fmt.Errorf("expected USER after CURRENT") - } - p.nextToken() - // Semicolon is optional for SHOW TOKEN - if p.curToken.Type == TokenSemicolon { + if p.curToken.Type == TokenUser { p.nextToken() + // Semicolon is optional for SHOW CURRENT USER + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return NewCommand("show_current_user"), nil + } else if p.curToken.Type == TokenModel { + p.nextToken() + // Semicolon is optional for SHOW CURRENT MODEL + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return NewCommand("show_current_model"), nil + } else { + return nil, fmt.Errorf("expected USER or MODEL after CURRENT") } - return NewCommand("show_current_user"), nil case TokenUser: return p.parseShowUser() case TokenRole: @@ -369,6 +380,8 @@ func (p *Parser) parseShowCommand() (*Command, error) { return p.parseShowProvider() case TokenModel: return p.parseShowModel() + case TokenInstance: + return p.parseShowInstance() default: return nil, fmt.Errorf("unknown SHOW target: %s", p.curToken.Value) } @@ -524,8 +537,6 @@ func (p *Parser) parseCreateCommand() (*Command, error) { return p.parseCreateRole() case TokenModel: return p.parseCreateModelProvider() - case TokenProvider: - return p.parseCreateProvider() case TokenDataset: return p.parseCreateDataset() case TokenChat: @@ -534,11 +545,23 @@ func (p *Parser) parseCreateCommand() (*Command, error) { return p.parseCreateToken() case TokenIndex: return p.parseCreateIndex() + case TokenProvider: + return p.parseCreateProviderInstance() default: return nil, fmt.Errorf("unknown CREATE target: %s", p.curToken.Value) } } +func (p *Parser) parseAddCommand() (*Command, error) { + p.nextToken() // consume ADD + switch p.curToken.Type { + case TokenProvider: + return p.parseAddProvider() + default: + return nil, fmt.Errorf("unknown ADD target: %s", p.curToken.Value) + } +} + func (p *Parser) parseCreateToken() (*Command, error) { p.nextToken() // consume TOKEN @@ -689,10 +712,10 @@ func (p *Parser) parseCreateModelProvider() (*Command, error) { return cmd, nil } -// parseCreateProvider parses CREATE PROVIDER commands -// CREATE PROVIDER -// CREATE PROVIDER -func (p *Parser) parseCreateProvider() (*Command, error) { +// parseAddProvider parses ADD PROVIDER commands +// ADD PROVIDER +// ADD PROVIDER +func (p *Parser) parseAddProvider() (*Command, error) { p.nextToken() // consume PROVIDER providerName, err := p.parseQuotedString() @@ -700,7 +723,7 @@ func (p *Parser) parseCreateProvider() (*Command, error) { return nil, fmt.Errorf("expected provider name: %w", err) } - cmd := NewCommand("create_provider") + cmd := NewCommand("add_provider") cmd.Params["provider_name"] = providerName p.nextToken() @@ -804,8 +827,6 @@ func (p *Parser) parseDropCommand() (*Command, error) { return p.parseDropRole() case TokenModel: return p.parseDropModelProvider() - case TokenProvider: - return p.parseDropProvider() case TokenDataset: return p.parseDropDataset() case TokenChat: @@ -814,6 +835,19 @@ func (p *Parser) parseDropCommand() (*Command, error) { return p.parseDropToken() case TokenIndex: return p.parseDropIndex() + case TokenInstance: + return p.parseDropInstance() + default: + return nil, fmt.Errorf("unknown DROP target: %s", p.curToken.Value) + } +} + +func (p *Parser) parseDeleteCommand() (*Command, error) { + p.nextToken() // consume DELETE + + switch p.curToken.Type { + case TokenProvider: + return p.parseDeleteProvider() default: return nil, fmt.Errorf("unknown DROP target: %s", p.curToken.Value) } @@ -949,8 +983,8 @@ func (p *Parser) parseDropModelProvider() (*Command, error) { return cmd, nil } -// parseDropProvider parses DROP PROVIDER command -func (p *Parser) parseDropProvider() (*Command, error) { +// parseDeleteProvider parses DELETE PROVIDER command +func (p *Parser) parseDeleteProvider() (*Command, error) { p.nextToken() // consume PROVIDER providerName, err := p.parseQuotedString() @@ -958,7 +992,7 @@ func (p *Parser) parseDropProvider() (*Command, error) { return nil, fmt.Errorf("expected provider name: %w", err) } - cmd := NewCommand("drop_provider") + cmd := NewCommand("delete_provider") cmd.Params["provider_name"] = providerName p.nextToken() @@ -1015,6 +1049,8 @@ func (p *Parser) parseAlterCommand() (*Command, error) { return p.parseAlterRole() case TokenProvider: return p.parseAlterProvider() + case TokenInstance: + return p.parseAlterInstance() default: return nil, fmt.Errorf("unknown ALTER target: %s", p.curToken.Value) } @@ -1176,6 +1212,189 @@ func (p *Parser) parseAlterProvider() (*Command, error) { return cmd, nil } +// parseCreateProviderInstance parses CREATE PROVIDER INSTANCE command +// instance_name cannot be "default" +func (p *Parser) parseCreateProviderInstance() (*Command, error) { + p.nextToken() // consume PROVIDER + + providerName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected provider name: %w", err) + } + + p.nextToken() + if p.curToken.Type != TokenInstance { + return nil, fmt.Errorf("expected INSTANCE after provider name") + } + p.nextToken() + + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected instance name: %w", err) + } + + // Check if instance_name is "default" + if instanceName == "default" { + return nil, fmt.Errorf("instance name cannot be 'default'") + } + + p.nextToken() + apiKey, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected API key: %w", err) + } + + cmd := NewCommand("create_provider_instance") + cmd.Params["provider_name"] = providerName + cmd.Params["instance_name"] = instanceName + cmd.Params["api_key"] = apiKey + + p.nextToken() + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + +// parseListInstances parses LIST INSTANCES FROM PROVIDER command +func (p *Parser) parseListInstances() (*Command, error) { + p.nextToken() // consume INSTANCES + + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + providerName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected provider name after FROM PROVIDER: %w", err) + } + + cmd := NewCommand("list_provider_instances") + cmd.Params["provider_name"] = providerName + + p.nextToken() + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + +// parseShowInstance parses SHOW INSTANCE FROM PROVIDER command +func (p *Parser) parseShowInstance() (*Command, error) { + p.nextToken() // consume INSTANCE + + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected instance name: %w", err) + } + + p.nextToken() + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + providerName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected provider name after FROM PROVIDER: %w", err) + } + + cmd := NewCommand("show_provider_instance") + cmd.Params["instance_name"] = instanceName + cmd.Params["provider_name"] = providerName + + p.nextToken() + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + +// parseAlterInstance parses ALTER INSTANCE NAME FROM PROVIDER command +func (p *Parser) parseAlterInstance() (*Command, error) { + p.nextToken() // consume INSTANCE + + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected instance name: %w", err) + } + + p.nextToken() + if p.curToken.Type != TokenName { + return nil, fmt.Errorf("expected NAME") + } + p.nextToken() + + newName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected new instance name: %w", err) + } + + p.nextToken() + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + if p.curToken.Type != TokenProvider { + return nil, fmt.Errorf("expected PROVIDER after FROM") + } + p.nextToken() + + providerName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected provider name after FROM PROVIDER: %w", err) + } + + cmd := NewCommand("alter_provider_instance") + cmd.Params["instance_name"] = instanceName + cmd.Params["new_name"] = newName + cmd.Params["provider_name"] = providerName + + p.nextToken() + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + +// parseDropInstance parses DROP INSTANCE FROM PROVIDER command +func (p *Parser) parseDropInstance() (*Command, error) { + p.nextToken() // consume INSTANCE + + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected instance name: %w", err) + } + + p.nextToken() + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + providerName, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected provider name after FROM PROVIDER: %w", err) + } + + cmd := NewCommand("drop_provider_instance") + cmd.Params["instance_name"] = instanceName + cmd.Params["provider_name"] = providerName + + p.nextToken() + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + func (p *Parser) parseGrantCommand() (*Command, error) { p.nextToken() // consume GRANT @@ -1662,6 +1881,215 @@ func (p *Parser) parseSearchCommand() (*Command, error) { return cmd, nil } +func (p *Parser) parseListModelsOfProvider() (*Command, error) { + if p.curToken.Type != TokenModels { + return nil, fmt.Errorf("expected MODELS") + } + + p.nextToken() + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + // Parse first quoted string (could be instance_name or provider_name) + firstName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + // Check if there's a second quoted string (provider_name) + // If so, format is: LIST MODELS FROM + // If not, format is: LIST MODELS FROM + if p.curToken.Type == TokenQuotedString { + // Two arguments: instance_name and provider_name + instanceName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + cmd := NewCommand("list_instance_models") + cmd.Params["instance_name"] = instanceName + cmd.Params["provider_name"] = firstName + p.nextToken() + // Semicolon is optional for UNSET TOKEN + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil + } + + // Only one argument: provider_name + cmd := NewCommand("list_provider_models") + cmd.Params["provider_name"] = firstName + // Semicolon is optional for UNSET TOKEN + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + return cmd, nil +} + +func (p *Parser) parseEnableCommand() (*Command, error) { + p.nextToken() // consume ENABLE + + if p.curToken.Type != TokenModel { + return nil, fmt.Errorf("expected MODEL") + } + p.nextToken() + + modelName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + modelProvider, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + modelInstance, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + // Semicolon is optional for UNSET TOKEN + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + + cmd := NewCommand("enable_model") + cmd.Params["model_name"] = modelName + cmd.Params["instance_name"] = modelInstance + cmd.Params["provider_name"] = modelProvider + return cmd, nil +} + +func (p *Parser) parseDisableCommand() (*Command, error) { + p.nextToken() // consume DISABLE + + if p.curToken.Type != TokenModel { + return nil, fmt.Errorf("expected MODEL") + } + p.nextToken() + + modelName, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + if p.curToken.Type != TokenFrom { + return nil, fmt.Errorf("expected FROM") + } + p.nextToken() + + modelProvider, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + modelInstance, err := p.parseQuotedString() + if err != nil { + return nil, err + } + p.nextToken() + + // Semicolon is optional for UNSET TOKEN + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + + cmd := NewCommand("disable_model") + cmd.Params["model_name"] = modelName + cmd.Params["instance_name"] = modelInstance + cmd.Params["provider_name"] = modelProvider + return cmd, nil +} + +func (p *Parser) parseChatCommand() (*Command, error) { + p.nextToken() // consume CHAT + + var modelName string + var message string + + // Check if we have a quoted string that looks like a model identifier (contains two slashes) + // Format: 'provider/instance/model' or just 'message' + if p.curToken.Type == TokenQuotedString { + firstArg := p.curToken.Value + + // Check if it looks like a model identifier (contains exactly 2 slashes) + slashCount := strings.Count(firstArg, "/") + if slashCount == 2 { + // This is likely a model identifier, expect another quoted string for message + modelName = firstArg + p.nextToken() + + // After model name, expect message + if p.curToken.Type != TokenQuotedString { + return nil, fmt.Errorf("expected message after model name") + } + message = p.curToken.Value + p.nextToken() + } else { + // This is just a message, use current model + message = firstArg + p.nextToken() + } + } else if p.curToken.Type == TokenIdentifier { + // Context engine style: chat + message = p.curToken.Value + p.nextToken() + } else { + return nil, fmt.Errorf("expected model name (quoted string) or message") + } + + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + + cmd := NewCommand("chat_to_model") + if modelName != "" { + cmd.Params["model_name"] = modelName + } + cmd.Params["message"] = message + return cmd, nil +} + +func (p *Parser) parseUseCommand() (*Command, error) { + p.nextToken() // consume USE + + if p.curToken.Type != TokenModel { + return nil, fmt.Errorf("expected MODEL after USE") + } + p.nextToken() // consume MODEL + + // Parse model identifier in format 'provider/instance/model' + modelIdentifier, err := p.parseQuotedString() + if err != nil { + return nil, fmt.Errorf("expected model identifier in format 'provider/instance/model': %w", err) + } + p.nextToken() + + // Semicolon is optional + if p.curToken.Type == TokenSemicolon { + p.nextToken() + } + + cmd := NewCommand("use_model") + cmd.Params["model_identifier"] = modelIdentifier + return cmd, nil +} + func (p *Parser) parseParseCommand() (*Command, error) { p.nextToken() // consume PARSE @@ -1788,6 +2216,7 @@ func (p *Parser) parseUserStatement() (*Command, error) { return p.parseInsertCommand() case TokenSearch: return p.parseSearchCommand() + default: return nil, fmt.Errorf("invalid user statement: %s", p.curToken.Value) } diff --git a/internal/dao/database.go b/internal/dao/database.go index e6418d37e3..429d2f5be1 100644 --- a/internal/dao/database.go +++ b/internal/dao/database.go @@ -147,6 +147,11 @@ func InitDB() error { &entity.EvaluationResult{}, &entity.TimeRecord{}, &entity.License{}, + &entity.TenantModelInstance{}, + &entity.TenantModel{}, + &entity.TenantModelGroupMapping{}, + &entity.TenantModelProvider{}, + &entity.TenantModelGroup{}, } for _, m := range models { diff --git a/internal/dao/tenant_model.go b/internal/dao/tenant_model.go new file mode 100644 index 0000000000..bb3b4f41ba --- /dev/null +++ b/internal/dao/tenant_model.go @@ -0,0 +1,67 @@ +// +// 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 dao + +import ( + "ragflow/internal/entity" +) + +// TenantModelDAO tenant model data access object +type TenantModelDAO struct{} + +// NewTenantModelDAO create tenant model DAO +func NewTenantModelDAO() *TenantModelDAO { + return &TenantModelDAO{} +} + +func (dao *TenantModelDAO) Create(instance *entity.TenantModel) error { + return DB.Create(instance).Error +} + +func (dao *TenantModelDAO) DeleteByModelID(modelID string) (int64, error) { + result := DB.Unscoped().Where("id = ?", modelID).Delete(&entity.TenantModel{}) + return result.RowsAffected, result.Error +} + +// GetByID get tenant model by primary key (id) +func (dao *TenantModelDAO) GetByID(id string) (*entity.TenantModel, error) { + var model entity.TenantModel + err := DB.Where("id = ?", id).First(&model).Error + if err != nil { + return nil, err + } + return &model, nil +} + +func (dao *TenantModelDAO) GetModelByProviderIDAndInstanceIDAndModelName(providerID, instanceID, modelName string) (*entity.TenantModel, error) { + var model entity.TenantModel + err := DB.Where("provider_id = ? AND instance_id = ? AND model_name = ?", providerID, instanceID, modelName).First(&model).Error + if err != nil { + return nil, err + } + return &model, nil +} + +// GetModelsByInstanceID get all models by instance ID +func (dao *TenantModelDAO) GetModelsByInstanceID(instanceID string) ([]*entity.TenantModel, error) { + var models []*entity.TenantModel + err := DB.Where("instance_id = ?", instanceID).Find(&models).Error + if err != nil { + return nil, err + } + return models, nil +} diff --git a/internal/dao/tenant_model_group.go b/internal/dao/tenant_model_group.go new file mode 100644 index 0000000000..e2d26982c9 --- /dev/null +++ b/internal/dao/tenant_model_group.go @@ -0,0 +1,39 @@ +// +// 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 dao + +import ( + "ragflow/internal/entity" +) + +// TenantModelGroupDAO tenant model group data access object +type TenantModelGroupDAO struct{} + +// NewTenantModelGroupDAO create tenant model group DAO +func NewTenantModelGroupDAO() *TenantModelGroupDAO { + return &TenantModelGroupDAO{} +} + +// GetByID get tenant model group by primary key (id) +func (dao *TenantModelGroupDAO) GetByID(id string) (*entity.TenantModelGroup, error) { + var group entity.TenantModelGroup + err := DB.Where("id = ?", id).First(&group).Error + if err != nil { + return nil, err + } + return &group, nil +} diff --git a/internal/dao/tenant_model_group_mapping.go b/internal/dao/tenant_model_group_mapping.go new file mode 100644 index 0000000000..c06270d275 --- /dev/null +++ b/internal/dao/tenant_model_group_mapping.go @@ -0,0 +1,39 @@ +// +// 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 dao + +import ( + "ragflow/internal/entity" +) + +// TenantModelGroupMappingDAO tenant model group mapping data access object +type TenantModelGroupMappingDAO struct{} + +// NewTenantModelGroupMappingDAO create tenant model group mapping DAO +func NewTenantModelGroupMappingDAO() *TenantModelGroupMappingDAO { + return &TenantModelGroupMappingDAO{} +} + +// GetByID get tenant model group mapping by composite primary key +func (dao *TenantModelGroupMappingDAO) GetByID(groupID, providerID, instanceID, modelID string) (*entity.TenantModelGroupMapping, error) { + var mapping entity.TenantModelGroupMapping + err := DB.Where("group_id = ? AND provider_id = ? AND instance_id = ? AND model_id = ?", groupID, providerID, instanceID, modelID).First(&mapping).Error + if err != nil { + return nil, err + } + return &mapping, nil +} diff --git a/internal/dao/tenant_model_instance.go b/internal/dao/tenant_model_instance.go new file mode 100644 index 0000000000..97eb4304e2 --- /dev/null +++ b/internal/dao/tenant_model_instance.go @@ -0,0 +1,66 @@ +// +// 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 dao + +import ( + "ragflow/internal/entity" +) + +// TenantModelInstanceDAO tenant model instance data access object +type TenantModelInstanceDAO struct{} + +// NewTenantModelInstanceDAO create tenant model instance DAO +func NewTenantModelInstanceDAO() *TenantModelInstanceDAO { + return &TenantModelInstanceDAO{} +} + +func (dao *TenantModelInstanceDAO) Create(instance *entity.TenantModelInstance) error { + return DB.Create(instance).Error +} + +func (dao *TenantModelInstanceDAO) GetAllInstancesByProviderID(providerID string) ([]*entity.TenantModelInstance, error) { + var instances []*entity.TenantModelInstance + err := DB.Where("provider_id = ?", providerID).Find(&instances).Error + if err != nil { + return nil, err + } + return instances, nil +} + +func (dao *TenantModelInstanceDAO) GetByProviderIDAndInstanceName(providerID, instanceName string) (*entity.TenantModelInstance, error) { + var instance entity.TenantModelInstance + err := DB.Where("provider_id = ? AND instance_name = ?", providerID, instanceName).First(&instance).Error + if err != nil { + return nil, err + } + return &instance, nil +} + +// GetByID get tenant model instance by primary key (id) +func (dao *TenantModelInstanceDAO) GetByID(id string) (*entity.TenantModelInstance, error) { + var instance entity.TenantModelInstance + err := DB.Where("id = ?", id).First(&instance).Error + if err != nil { + return nil, err + } + return &instance, nil +} + +func (dao *TenantModelInstanceDAO) DeleteByProviderIDAndInstanceName(providerID, instanceName string) (int64, error) { + result := DB.Unscoped().Where("provider_id = ? and instance_name = ?", providerID, instanceName).Delete(&entity.TenantModelInstance{}) + return result.RowsAffected, result.Error +} diff --git a/internal/dao/tenant_model_provider.go b/internal/dao/tenant_model_provider.go new file mode 100644 index 0000000000..fd75353bdb --- /dev/null +++ b/internal/dao/tenant_model_provider.go @@ -0,0 +1,74 @@ +// +// 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 dao + +import ( + "ragflow/internal/entity" +) + +// TenantModelProviderDAO tenant model provider data access object +type TenantModelProviderDAO struct{} + +// NewTenantModelProviderDAO create tenant model provider DAO +func NewTenantModelProviderDAO() *TenantModelProviderDAO { + return &TenantModelProviderDAO{} +} + +func (dao *TenantModelProviderDAO) Create(provider *entity.TenantModelProvider) error { + return DB.Create(provider).Error +} + +// GetByID get tenant model provider by primary key (id) +func (dao *TenantModelProviderDAO) GetByID(id string) (*entity.TenantModelProvider, error) { + var provider entity.TenantModelProvider + err := DB.Where("id = ?", id).First(&provider).Error + if err != nil { + return nil, err + } + return &provider, nil +} + +// GetByTenantIDAndProviderName get the providers by tenant ID and provider name +func (dao *TenantModelProviderDAO) GetByTenantIDAndProviderName(tenantID, providerName string) (*entity.TenantModelProvider, error) { + var provider entity.TenantModelProvider + err := DB.Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).First(&provider).Error + if err != nil { + return nil, err + } + return &provider, nil +} + +// DeleteByTenantID deletes all model providers by tenant ID (hard delete) +func (dao *TenantModelProviderDAO) DeleteByTenantID(tenantID string) (int64, error) { + result := DB.Unscoped().Where("tenant_id = ?", tenantID).Delete(&entity.TenantModelProvider{}) + return result.RowsAffected, result.Error +} + +// DeleteByTenantID deletes all providers by tenant ID (hard delete) +func (dao *TenantModelProviderDAO) DeleteByTenantIDAndProviderName(tenantID, providerName string) (int64, error) { + result := DB.Unscoped().Where("tenant_id = ? AND provider_name = ?", tenantID, providerName).Delete(&entity.TenantModelProvider{}) + return result.RowsAffected, result.Error +} + +// ListByID list tenant model providers by ID +func (dao *TenantModelProviderDAO) ListByID(id string) ([]string, error) { + var providerNames []string + err := DB.Model(&entity.TenantModelProvider{}). + Where("tenant_id = ?", id). + Pluck("provider_name", &providerNames).Error + return providerNames, err +} diff --git a/internal/entity/model.go b/internal/entity/model.go index a25ba026d0..0b8f208cb3 100644 --- a/internal/entity/model.go +++ b/internal/entity/model.go @@ -21,6 +21,7 @@ import ( "fmt" "os" "path/filepath" + "ragflow/internal/entity/models" "strings" ) @@ -135,18 +136,21 @@ type Features struct { // Model represents a single LLM model type Model struct { - Name string `json:"name"` - MaxTokens int `json:"max_tokens"` - ModelTypes []string `json:"model_types"` - Features Features `json:"features"` + Name string `json:"name"` + MaxTokens int `json:"max_tokens"` + ModelTypes []string `json:"model_types"` + Features Features `json:"features"` + ModelTypeMap map[string]bool } // Provider represents an LLM provider type Provider struct { - Name string `json:"name"` - Tags string `json:"tags"` - URL string `json:"url"` - Models []Model `json:"models"` + Name string `json:"name"` + Tags string `json:"tags"` + URL string `json:"url"` + URLSuffix models.URLSuffix `json:"url_suffix"` + Models []Model `json:"models"` + ModelDriver models.ModelDriver } // ProviderManager manages provider and model operations @@ -171,6 +175,8 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) { return nil, fmt.Errorf("error reading directory %s: %w", dirPath, err) } + modelFactory := models.NewModelFactory() + // Iterate through all files for _, file := range files { // Skip directories @@ -187,7 +193,8 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) { filePath := filepath.Join(dirPath, file.Name()) // Read the file - data, err := os.ReadFile(filePath) + var data []byte + data, err = os.ReadFile(filePath) if err != nil { return nil, fmt.Errorf("error reading file %s: %w", filePath, err) } @@ -198,6 +205,18 @@ func NewProviderManager(dirPath string) (*ProviderManager, error) { return nil, fmt.Errorf("error parsing JSON from file %s: %w", filePath, err) } + for _, model := range provider.Models { + model.ModelTypeMap = make(map[string]bool) + for _, modelType := range model.ModelTypes { + model.ModelTypeMap[modelType] = true + } + } + + provider.ModelDriver, err = modelFactory.CreateModelDriver(provider.Name, provider.URL, provider.URLSuffix) + if err != nil { + return nil, fmt.Errorf("error creating model driver for provider %s: %w", provider.Name, err) + } + // Add to providers list providers = append(providers, provider) } @@ -218,9 +237,10 @@ func (pm *ProviderManager) ListProviders() ([]map[string]interface{}, error) { for _, provider := range pm.Providers { providerData := map[string]interface{}{ - "name": provider.Name, - "tags": provider.Tags, - "url": provider.URL, + "name": provider.Name, + "tags": provider.Tags, + "url": provider.URL, + "url_suffix": provider.URLSuffix, } providers = append(providers, providerData) } @@ -235,7 +255,7 @@ func (pm *ProviderManager) ListProviders() ([]map[string]interface{}, error) { // 2. Show specific provider information (including base_url) func (pm *ProviderManager) GetProviderByName(providerName string) (map[string]interface{}, error) { - provider := pm.findProvider(providerName) + provider := pm.FindProvider(providerName) if provider == nil { return nil, fmt.Errorf("provider '%s' not found", providerName) } @@ -252,7 +272,7 @@ func (pm *ProviderManager) GetProviderByName(providerName string) (map[string]in // 3. List models under a specific provider func (pm *ProviderManager) ListModels(providerName string) ([]map[string]interface{}, error) { - provider := pm.findProvider(providerName) + provider := pm.FindProvider(providerName) if provider == nil { return nil, fmt.Errorf("provider '%s' not found", providerName) } @@ -276,7 +296,7 @@ func (pm *ProviderManager) ListModels(providerName string) ([]map[string]interfa } func (pm *ProviderManager) GetModelByName(providerName, modelName string) (*Model, error) { - provider := pm.findProvider(providerName) + provider := pm.FindProvider(providerName) if provider == nil { return nil, fmt.Errorf("provider '%s' not found", providerName) } @@ -287,6 +307,39 @@ func (pm *ProviderManager) GetModelByName(providerName, modelName string) (*Mode return model, nil } +func (pm *ProviderManager) GetModelUrl(providerName, modelName, modelType string) (*string, *string, error) { + provider := pm.FindProvider(providerName) + if provider == nil { + return nil, nil, fmt.Errorf("provider '%s' not found", providerName) + } + model := pm.findModel(provider, modelName) + if model == nil { + return nil, nil, fmt.Errorf("model '%s' not found", modelName) + } + + if !model.ModelTypeMap[modelType] { + return nil, nil, fmt.Errorf("model '%s' does not support model type '%s'", modelName, modelType) + } + + switch modelType { + case "chat": + url := fmt.Sprintf("%s%s", provider.URL, provider.URLSuffix.Chat) + return &url, nil, nil + case "async_chat": + chatUrl := fmt.Sprintf("%s%s", provider.URL, provider.URLSuffix.AsyncChat) + resultUrl := fmt.Sprintf("%s%s", provider.URL, provider.URLSuffix.AsyncResult) + return &chatUrl, &resultUrl, nil + case "embedding": + url := fmt.Sprintf("%s%s", provider.URL, provider.URLSuffix.Embedding) + return &url, nil, nil + case "rerank": + url := fmt.Sprintf("%s%s", provider.URL, provider.URLSuffix.Rerank) + return &url, nil, nil + default: + return nil, nil, fmt.Errorf("model '%s' does not support model type '%s'", modelName, modelType) + } +} + // 4. Search specific model information with filtering by max_tokens or type func (pm *ProviderManager) SearchModelInfo(providerName, modelName string, filterBy string, filterValue interface{}) ModelResponse { resp := ModelResponse{ @@ -295,7 +348,7 @@ func (pm *ProviderManager) SearchModelInfo(providerName, modelName string, filte Message: "success", } - provider := pm.findProvider(providerName) + provider := pm.FindProvider(providerName) if provider == nil { resp.Code = 404 resp.Message = fmt.Sprintf("Provider '%s' not found", providerName) @@ -481,7 +534,7 @@ func modelHasFeature(features Features, featureType string) bool { } // Helper: Find provider by name -func (pm *ProviderManager) findProvider(name string) *Provider { +func (pm *ProviderManager) FindProvider(name string) *Provider { for i := range pm.Providers { if strings.EqualFold(pm.Providers[i].Name, name) { return &pm.Providers[i] diff --git a/internal/entity/models/dummy.go b/internal/entity/models/dummy.go new file mode 100644 index 0000000000..84ebea3191 --- /dev/null +++ b/internal/entity/models/dummy.go @@ -0,0 +1,50 @@ +// +// 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 ( + "fmt" +) + +// DummyModel implements ModelDriver for Zhipu AI (智谱 AI) +type DummyModel struct { + BaseURL string + URLSuffix URLSuffix +} + +// NewDummyModel creates a new Zhipu AI model instance +func NewDummyModel(baseURL string, urlSuffix URLSuffix) *DummyModel { + return &DummyModel{ + BaseURL: baseURL, + URLSuffix: urlSuffix, + } +} + +// Chat sends a message and returns response +func (z *DummyModel) Chat(modelName, apiKey, message *string, genConf map[string]interface{}) (string, error) { + return "", fmt.Errorf("not implemented") +} + +// ChatStreamly sends a message and streams response +func (z *DummyModel) ChatStreamly(modelName, apiKey, message *string, genConf map[string]interface{}) (<-chan string, error) { + return nil, fmt.Errorf("not implemented") +} + +// EncodeToEmbedding encodes a list of texts into embeddings +func (z *DummyModel) EncodeToEmbedding(modelName, apiKey *string, texts []string) ([][]float64, error) { + return nil, fmt.Errorf("not implemented") +} diff --git a/internal/entity/models/factory.go b/internal/entity/models/factory.go new file mode 100644 index 0000000000..2531c50f23 --- /dev/null +++ b/internal/entity/models/factory.go @@ -0,0 +1,41 @@ +// +// 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 ( + "strings" +) + +// ModelFactory creates ModelDriver instances based on provider name +type ModelFactory struct { +} + +// NewModelFactory creates a new ModelFactory +func NewModelFactory() *ModelFactory { + return &ModelFactory{} +} + +// CreateModelDriver creates a ModelDriver for the given provider and model +func (f *ModelFactory) CreateModelDriver(providerName string, baseURL string, urlSuffix URLSuffix) (ModelDriver, error) { + providerLower := strings.ToLower(providerName) + switch providerLower { + case "zhipu-ai": + return NewZhipuAIModel(baseURL, urlSuffix), nil + default: + return NewDummyModel(baseURL, urlSuffix), nil + } +} diff --git a/internal/entity/models/types.go b/internal/entity/models/types.go new file mode 100644 index 0000000000..42f80039ee --- /dev/null +++ b/internal/entity/models/types.go @@ -0,0 +1,20 @@ +package models + +// EmbeddingModel interface for embedding models +type ModelDriver interface { + // Chat sends a message and returns response + Chat(modelName, apiKey, message *string, genConf map[string]interface{}) (string, error) + // ChatStreamly sends a message and streams response + ChatStreamly(modelName, apiKey, message *string, genConf map[string]interface{}) (<-chan string, error) + // Encode encodes a list of texts into embeddings + EncodeToEmbedding(modelName, apiKey *string, texts []string) ([][]float64, error) +} + +// URLSuffix represents the URL suffixes for different API endpoints +type URLSuffix struct { + Chat string `json:"chat"` + AsyncChat string `json:"async_chat"` + AsyncResult string `json:"async_result"` + Embedding string `json:"embedding"` + Rerank string `json:"rerank"` +} diff --git a/internal/entity/models/zhipu-ai.go b/internal/entity/models/zhipu-ai.go new file mode 100644 index 0000000000..1f17c8f322 --- /dev/null +++ b/internal/entity/models/zhipu-ai.go @@ -0,0 +1,307 @@ +// +// 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" +) + +// ZhipuAIModel implements ModelDriver for Zhipu AI (智谱 AI) +type ZhipuAIModel struct { + BaseURL string + URLSuffix URLSuffix +} + +// NewZhipuAIModel creates a new Zhipu AI model instance +func NewZhipuAIModel(baseURL string, urlSuffix URLSuffix) *ZhipuAIModel { + return &ZhipuAIModel{ + BaseURL: baseURL, + URLSuffix: urlSuffix, + } +} + +// Chat sends a message and returns response +func (z *ZhipuAIModel) Chat(modelName, apiKey, message *string, genConf map[string]interface{}) (string, error) { + if message == nil { + return "", fmt.Errorf("message is nil") + } + + url := fmt.Sprintf("%s/%s", z.BaseURL, z.URLSuffix.Chat) + + // Build request body + reqBody := map[string]interface{}{ + "model": modelName, + "messages": []map[string]string{ + {"role": "user", "content": *message}, + }, + "stream": false, + "temperature": 1, + } + + // Add generation config if provided + if genConf != nil { + if maxTokens, ok := genConf["max_tokens"]; ok { + reqBody["max_tokens"] = maxTokens + } + if temperature, ok := genConf["temperature"]; ok { + reqBody["temperature"] = temperature + } + if topP, ok := genConf["top_p"]; ok { + reqBody["top_p"] = topP + } + } + + jsonData, err := json.Marshal(reqBody) + if err != nil { + return "", fmt.Errorf("failed to marshal request: %w", err) + } + + req, err := http.NewRequest("POST", url, bytes.NewBuffer(jsonData)) + if err != nil { + return "", fmt.Errorf("failed to create request: %w", err) + } + + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", *apiKey)) + + client := &http.Client{} + resp, err := client.Do(req) + if err != nil { + return "", fmt.Errorf("failed to send request: %w", err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + return "", fmt.Errorf("failed to read response: %w", err) + } + + if resp.StatusCode != http.StatusOK { + return "", 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 "", fmt.Errorf("failed to parse response: %w", err) + } + + choices, ok := result["choices"].([]interface{}) + if !ok || len(choices) == 0 { + return "", fmt.Errorf("no choices in response") + } + + firstChoice, ok := choices[0].(map[string]interface{}) + if !ok { + return "", fmt.Errorf("invalid choice format") + } + + messageMap, ok := firstChoice["message"].(map[string]interface{}) + if !ok { + return "", fmt.Errorf("invalid message format") + } + + content, ok := messageMap["content"].(string) + if !ok { + return "", fmt.Errorf("invalid content format") + } + + return content, nil +} + +// ChatStreamly sends a message and streams response +func (z *ZhipuAIModel) ChatStreamly(modelName, apiKey, message *string, genConf map[string]interface{}) (<-chan string, error) { + url := fmt.Sprintf("%s/chat/completions", z.BaseURL) + + // Build request body with streaming enabled + reqBody := map[string]interface{}{ + "model": modelName, + "messages": []map[string]string{ + {"role": "user", "content": *message}, + }, + "stream": true, + } + + // Add generation config if provided + if genConf != nil { + if maxTokens, ok := genConf["max_tokens"]; ok { + reqBody["max_tokens"] = maxTokens + } + if temperature, ok := genConf["temperature"]; ok { + reqBody["temperature"] = temperature + } + if topP, ok := genConf["top_p"]; ok { + reqBody["top_p"] = topP + } + } + + 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", apiKey)) + + client := &http.Client{} + resp, err := client.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return nil, fmt.Errorf("API request failed with status %d: %s", resp.StatusCode, string(body)) + } + + // Create channel for streaming + resultChan := make(chan string) + + go func() { + defer close(resultChan) + defer resp.Body.Close() + + decoder := json.NewDecoder(resp.Body) + for { + var event map[string]interface{} + if err := decoder.Decode(&event); err != nil { + if err == io.EOF { + break + } + return + } + + choices, ok := event["choices"].([]interface{}) + if !ok || len(choices) == 0 { + continue + } + + firstChoice, ok := choices[0].(map[string]interface{}) + if !ok { + continue + } + + delta, ok := firstChoice["delta"].(map[string]interface{}) + if !ok { + continue + } + + content, ok := delta["content"].(string) + if ok && content != "" { + resultChan <- content + } + + finishReason, ok := firstChoice["finish_reason"].(string) + if ok && finishReason != "" { + break + } + } + }() + + return resultChan, nil +} + +// EncodeToEmbedding encodes a list of texts into embeddings +func (z *ZhipuAIModel) EncodeToEmbedding(modelName, apiKey *string, texts []string) ([][]float64, error) { + url := fmt.Sprintf("%s/embedding", z.BaseURL) + + embeddings := make([][]float64, len(texts)) + + for i, text := range texts { + reqBody := map[string]interface{}{ + "model": modelName, + "input": text, + } + + 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", apiKey)) + + client := &http.Client{} + resp, err := client.Do(req) + if err != nil { + return nil, fmt.Errorf("failed to send request: %w", err) + } + + body, err := io.ReadAll(resp.Body) + resp.Body.Close() + + 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) + } + + data, ok := result["data"].([]interface{}) + if !ok || len(data) == 0 { + return nil, fmt.Errorf("no data in response") + } + + firstData, ok := data[0].(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("invalid data format") + } + + embeddingSlice, ok := firstData["embedding"].([]interface{}) + if !ok { + return nil, fmt.Errorf("invalid embedding format") + } + + embedding := make([]float64, len(embeddingSlice)) + for j, v := range embeddingSlice { + switch val := v.(type) { + case float64: + embedding[j] = val + case float32: + embedding[j] = float64(val) + default: + return nil, fmt.Errorf("unexpected embedding value type") + } + } + + embeddings[i] = embedding + } + + return embeddings, nil +} diff --git a/internal/entity/tenant_model.go b/internal/entity/tenant_model.go new file mode 100644 index 0000000000..72e4b41a5a --- /dev/null +++ b/internal/entity/tenant_model.go @@ -0,0 +1,33 @@ +// +// 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 entity + +// TenantModel tenant model table +type TenantModel struct { + ID string `gorm:"column:id;primaryKey;size:32" json:"id"` + ModelName string `gorm:"column:model_name;size:128" json:"model_name"` + ProviderID string `gorm:"column:provider_id;size:32;not null" json:"provider_id"` + InstanceID string `gorm:"column:instance_id;size:32;not null;index" json:"instance_id"` + ModelType string `gorm:"column:model_type;size:32;not null" json:"model_type"` + Status string `gorm:"column:status;size:32;default:'active'" json:"status"` + BaseModel +} + +// TableName specify table name +func (TenantModel) TableName() string { + return "tenant_model" +} diff --git a/internal/entity/tenant_model_group.go b/internal/entity/tenant_model_group.go new file mode 100644 index 0000000000..9e16bc6cbe --- /dev/null +++ b/internal/entity/tenant_model_group.go @@ -0,0 +1,31 @@ +// +// 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 entity + +// TenantModelGroup tenant model group table +type TenantModelGroup struct { + ID string `gorm:"column:id;primaryKey;size:32" json:"id"` + GroupType string `gorm:"column:group_type;size:32;not null" json:"group_type"` + ModelName *string `gorm:"column:model_name;size:128" json:"model_name,omitempty"` + Strategy string `gorm:"column:strategy;size:32;default:'weighted'" json:"strategy"` + BaseModel +} + +// TableName specify table name +func (TenantModelGroup) TableName() string { + return "tenant_model_group" +} diff --git a/internal/entity/tenant_model_group_mapping.go b/internal/entity/tenant_model_group_mapping.go new file mode 100644 index 0000000000..b7e6f9d504 --- /dev/null +++ b/internal/entity/tenant_model_group_mapping.go @@ -0,0 +1,33 @@ +// +// 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 entity + +// TenantModelGroupMapping tenant model group mapping table +type TenantModelGroupMapping struct { + GroupID string `gorm:"column:group_id;primaryKey;size:32;index" json:"group_id"` + ProviderID string `gorm:"column:provider_id;primaryKey;size:32" json:"provider_id"` + InstanceID string `gorm:"column:instance_id;primaryKey;size:32" json:"instance_id"` + ModelID string `gorm:"column:model_id;primaryKey;size:32;index" json:"model_id"` + Weight int `gorm:"column:weight;default:100" json:"weight"` + Status string `gorm:"column:status;size:32;default:'active'" json:"status"` + BaseModel +} + +// TableName specify table name +func (TenantModelGroupMapping) TableName() string { + return "tenant_model_group_mapping" +} diff --git a/internal/entity/tenant_model_instance.go b/internal/entity/tenant_model_instance.go new file mode 100644 index 0000000000..de5da075af --- /dev/null +++ b/internal/entity/tenant_model_instance.go @@ -0,0 +1,32 @@ +// +// 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 entity + +// TenantModelInstance tenant model instance table +type TenantModelInstance struct { + ID string `gorm:"column:id;primaryKey;size:32" json:"id"` + InstanceName string `gorm:"column:instance_name;size:128;not null" json:"instance_name"` + ProviderID string `gorm:"column:provider_id;size:32;not null;index" json:"provider_id"` + APIKey string `gorm:"column:api_key;size:512;not null;uniqueIndex" json:"api_key"` + Status string `gorm:"column:status;size:32;default:'active'" json:"status"` + BaseModel +} + +// TableName specify table name +func (TenantModelInstance) TableName() string { + return "tenant_model_instance" +} diff --git a/internal/entity/tenant_model_provider.go b/internal/entity/tenant_model_provider.go new file mode 100644 index 0000000000..db65188359 --- /dev/null +++ b/internal/entity/tenant_model_provider.go @@ -0,0 +1,30 @@ +// +// 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 entity + +// TenantModelProvider tenant model provider table +type TenantModelProvider struct { + ID string `gorm:"column:id;primaryKey;size:32" json:"id"` + ProviderName string `gorm:"column:provider_name;size:128;not null;index:idx_tenant_provider_unique,unique" json:"provider_name"` + TenantID string `gorm:"column:tenant_id;size:32;not null;index;index:idx_tenant_provider_unique,unique" json:"tenant_id"` + BaseModel +} + +// TableName specify table name +func (TenantModelProvider) TableName() string { + return "tenant_model_provider" +} diff --git a/internal/handler/providers.go b/internal/handler/providers.go index a936bcd450..be99355555 100644 --- a/internal/handler/providers.go +++ b/internal/handler/providers.go @@ -28,13 +28,17 @@ import ( // ProviderHandler provider handler type ProviderHandler struct { - userService *service.UserService + userService *service.UserService + modelProviderService *service.ModelProviderService + userTenantDAO *dao.UserTenantDAO } // NewProviderHandler create provider handler -func NewProviderHandler(userService *service.UserService) *ProviderHandler { +func NewProviderHandler(userService *service.UserService, modelProviderService *service.ModelProviderService) *ProviderHandler { return &ProviderHandler{ - userService: userService, + userService: userService, + modelProviderService: modelProviderService, + userTenantDAO: dao.NewUserTenantDAO(), } } @@ -58,12 +62,98 @@ func (h *ProviderHandler) ListProviders(c *gin.Context) { return } + for _, provider := range providers { + delete(provider, "url_suffix") + delete(provider, "tags") + } + c.JSON(http.StatusOK, gin.H{ "code": 0, "message": "success", "data": providers, }) + return } + + userID := c.GetString("user_id") + + // list tenant providers + providers, errorCode, err := h.modelProviderService.ListProvidersOfTenant(userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + "data": nil, + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + "data": providers, + }) + return +} + +type AddProviderRequest struct { + ProviderName string `json:"provider_name" binding:"required"` +} + +func (h *ProviderHandler) AddProvider(c *gin.Context) { + + var req AddProviderRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": err.Error(), + "data": false, + }) + return + } + + userID := c.GetString("user_id") + + errorCode, err := h.modelProviderService.AddModelProvider(req.ProviderName, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + }) +} + +func (h *ProviderHandler) DeleteProvider(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + userID := c.GetString("user_id") + + errorCode, err := h.modelProviderService.DeleteModelProvider(providerName, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + }) } func (h *ProviderHandler) ShowProvider(c *gin.Context) { @@ -146,3 +236,352 @@ func (h *ProviderHandler) ShowModel(c *gin.Context) { "data": model, }) } + +type CreateProviderInstanceRequest struct { + InstanceName string `json:"instance_name" binding:"required"` + APIKey string `json:"api_key" binding:"required"` +} + +func (h *ProviderHandler) CreateProviderInstance(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + var req CreateProviderInstanceRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": err.Error(), + }) + return + } + + // Check if instance name is "default" + if req.InstanceName == "default" { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": "Instance name cannot be 'default'", + }) + return + } + + userID := c.GetString("user_id") + + _, err := h.modelProviderService.CreateProviderInstance(providerName, req.InstanceName, req.APIKey, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeServerError, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + }) +} + +func (h *ProviderHandler) ListProviderInstances(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + userID := c.GetString("user_id") + + instances, errorCode, err := h.modelProviderService.ListProviderInstances(providerName, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + "data": instances, + }) +} + +func (h *ProviderHandler) ShowProviderInstance(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + userID := c.GetString("user_id") + + // Get tenant ID from user + instance, errorCode, err := h.modelProviderService.ShowProviderInstance(providerName, instanceName, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + "data": instance, + }) +} + +type AlterProviderInstanceRequest struct { + LLMName string `json:"llm_name" binding:"required"` +} + +func (h *ProviderHandler) AlterProviderInstance(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + var req AlterProviderInstanceRequest + if err := c.ShouldBindJSON(&req); err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeBadRequest, + "message": err.Error(), + }) + return + } + + userID := c.GetString("user_id") + if userID == "" { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeUnauthorized, + "message": "Unauthorized", + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeNotFound, + "message": "success", + }) +} + +func (h *ProviderHandler) DropProviderInstance(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + userID := c.GetString("user_id") + + _, err := h.modelProviderService.DropProviderInstance(providerName, instanceName, userID) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeServerError, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + }) +} + +func (h *ProviderHandler) ListInstanceModels(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + models, err := h.modelProviderService.ListInstanceModels(providerName, instanceName, c.GetString("user_id")) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeNotFound, + "message": err.Error(), + }) + return + } + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + "data": models, + }) +} + +type EnableOrDisableModelRequest struct { + Status string `json:"status" binding:"required"` +} + +func (h *ProviderHandler) EnableOrDisableModel(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + modelName := c.Param("model_name") + if modelName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model name is required", + }) + return + } + + var req EnableOrDisableModelRequest + 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 + } + + userID := c.GetString("user_id") + + _, err := h.modelProviderService.UpdateModelStatus(providerName, instanceName, modelName, userID, req.Status) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": common.CodeServerError, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": "success", + }) +} + +type ChatToModelRequest struct { + Message string `json:"message" binding:"required"` +} + +func (h *ProviderHandler) ChatToModel(c *gin.Context) { + providerName := c.Param("provider_name") + if providerName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Provider name is required", + }) + return + } + + instanceName := c.Param("instance_name") + if instanceName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Instance name is required", + }) + return + } + + modelName := c.Param("model_name") + if modelName == "" { + c.JSON(http.StatusBadRequest, gin.H{ + "code": 400, + "message": "Model name is required", + }) + return + } + + var req ChatToModelRequest + 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 + } + + userID := c.GetString("user_id") + + response, errorCode, err := h.modelProviderService.ChatToModel(providerName, instanceName, modelName, userID, req.Message) + if err != nil { + c.JSON(http.StatusOK, gin.H{ + "code": errorCode, + "message": err.Error(), + }) + return + } + + c.JSON(http.StatusOK, gin.H{ + "code": 0, + "message": response, + }) +} diff --git a/internal/router/router.go b/internal/router/router.go index 3aecefe765..72f485d25a 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -193,9 +193,19 @@ func (r *Router) Setup(engine *gin.Engine) { provider := v1.Group("/providers") { provider.GET("/", r.providerHandler.ListProviders) + provider.POST("/", r.providerHandler.AddProvider) provider.GET("/:provider_name", r.providerHandler.ShowProvider) + provider.DELETE("/:provider_name", r.providerHandler.DeleteProvider) provider.GET("/:provider_name/models", r.providerHandler.ListModels) provider.GET("/:provider_name/models/:model_name", r.providerHandler.ShowModel) + provider.POST("/:provider_name/instances", r.providerHandler.CreateProviderInstance) + provider.GET("/:provider_name/instances", r.providerHandler.ListProviderInstances) + provider.GET("/:provider_name/instances/:instance_name", r.providerHandler.ShowProviderInstance) + provider.PUT("/:provider_name/instances/:instance_name", r.providerHandler.AlterProviderInstance) + provider.DELETE("/:provider_name/instances/:instance_name", r.providerHandler.DropProviderInstance) + provider.GET("/:provider_name/instances/:instance_name/models", r.providerHandler.ListInstanceModels) + provider.PUT("/:provider_name/instances/:instance_name/models/:model_name", r.providerHandler.EnableOrDisableModel) + provider.POST("/:provider_name/instances/:instance_name/models/:model_name", r.providerHandler.ChatToModel) } } diff --git a/internal/service/kb.go b/internal/service/kb.go index 05cf9ee2bb..bf5d2f72f2 100644 --- a/internal/service/kb.go +++ b/internal/service/kb.go @@ -184,7 +184,7 @@ func (s *KnowledgebaseService) CreateKB(req *CreateKBRequest, tenantID string) ( } // Create in database - if err := s.kbDAO.Create(kb); err != nil { + if err = s.kbDAO.Create(kb); err != nil { return nil, common.CodeServerError, fmt.Errorf("failed to create knowledge base: %w", err) } diff --git a/internal/service/model_service.go b/internal/service/model_service.go index 1cc4428d84..50e53aa97f 100644 --- a/internal/service/model_service.go +++ b/internal/service/model_service.go @@ -18,8 +18,10 @@ package service import ( "context" + "errors" "fmt" "net/http" + "ragflow/internal/common" "ragflow/internal/dao" "ragflow/internal/entity" "strings" @@ -115,3 +117,428 @@ func (p *ModelProviderImpl) GetRerankModel(ctx context.Context, tenantID string, // TODO: implement rerank model creation return nil, fmt.Errorf("rerank model not implemented yet for model: %s", compositeModelName) } + +func NewModelProviderService() *ModelProviderService { + return &ModelProviderService{ + modelProviderDAO: dao.NewTenantModelProviderDAO(), + modelInstanceDAO: dao.NewTenantModelInstanceDAO(), + modelDAO: dao.NewTenantModelDAO(), + modelGroupDAO: dao.NewTenantModelGroupDAO(), + modelGroupMappingDAO: dao.NewTenantModelGroupMappingDAO(), + userTenantDAO: dao.NewUserTenantDAO(), + } +} + +type ModelProviderService struct { + modelProviderDAO *dao.TenantModelProviderDAO + modelInstanceDAO *dao.TenantModelInstanceDAO + modelDAO *dao.TenantModelDAO + modelGroupDAO *dao.TenantModelGroupDAO + modelGroupMappingDAO *dao.TenantModelGroupMappingDAO + userTenantDAO *dao.UserTenantDAO +} + +func (m *ModelProviderService) AddModelProvider(providerName, userID string) (common.ErrorCode, error) { + + _, err := dao.GetModelProviderManager().GetProviderByName(providerName) + if err != nil { + return common.CodeNotFound, err + } + + 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 + + providerID, err := generateUUID1Hex() + if err != nil { + return common.CodeServerError, errors.New("fail to get UUID") + } + + now := time.Now().Unix() + nowDate := time.Now().Truncate(time.Second) + tenantModelProvider := &entity.TenantModelProvider{ + ID: providerID, + ProviderName: providerName, + TenantID: tenantID, + } + tenantModelProvider.CreateTime = &now + tenantModelProvider.UpdateTime = &now + tenantModelProvider.CreateDate = &nowDate + tenantModelProvider.UpdateDate = &nowDate + err = m.modelProviderDAO.Create(tenantModelProvider) + if err != nil { + return common.CodeServerError, errors.New("fail to create model provider") + } + return common.CodeSuccess, nil +} + +func (m *ModelProviderService) ListProvidersOfTenant(userID string) ([]map[string]interface{}, common.ErrorCode, error) { + + 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 + + providerNames, err := m.modelProviderDAO.ListByID(tenantID) + if err != nil { + return nil, common.CodeServerError, err + } + + var result []map[string]interface{} + for _, providerName := range providerNames { + provider, err := dao.GetModelProviderManager().GetProviderByName(providerName) + if err != nil { + return nil, common.CodeServerError, err + } + result = append(result, provider) + } + + return result, common.CodeSuccess, nil +} + +func (m *ModelProviderService) DeleteModelProvider(providerName, userID string) (common.ErrorCode, error) { + 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 + + _, err = m.modelProviderDAO.DeleteByTenantIDAndProviderName(tenantID, providerName) + if err != nil { + return common.CodeServerError, err + } + + return common.CodeSuccess, nil +} + +func (m *ModelProviderService) CreateProviderInstance(providerName, instanceName, apiKey, 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, providerName) + if err != nil { + return common.CodeServerError, err + } + + instanceID, err := generateUUID1Hex() + if err != nil { + return common.CodeServerError, errors.New("fail to get UUID") + } + + now := time.Now().Unix() + nowDate := time.Now().Truncate(time.Second) + tenantModelProvider := &entity.TenantModelInstance{ + ID: instanceID, + InstanceName: instanceName, + ProviderID: provider.ID, + APIKey: apiKey, + Status: "active", + } + tenantModelProvider.CreateTime = &now + tenantModelProvider.UpdateTime = &now + tenantModelProvider.CreateDate = &nowDate + tenantModelProvider.UpdateDate = &nowDate + err = m.modelInstanceDAO.Create(tenantModelProvider) + + if err != nil { + return common.CodeServerError, errors.New("fail to create model provider") + } + return common.CodeSuccess, nil +} + +func (m *ModelProviderService) ListProviderInstances(providerName, userID string) ([]map[string]interface{}, common.ErrorCode, error) { + + // 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 + } + + // Check if provider exists + instances, err := m.modelInstanceDAO.GetAllInstancesByProviderID(provider.ID) + if err != nil { + return nil, common.CodeServerError, err + } + + var result []map[string]interface{} + for _, instance := range instances { + result = append(result, map[string]interface{}{ + "id": instance.ID, + "instanceName": instance.InstanceName, + "providerID": instance.ProviderID, + "apiKey": instance.APIKey, + "status": instance.Status, + }) + } + + return result, common.CodeSuccess, nil +} + +func (m *ModelProviderService) ShowProviderInstance(providerName, instanceName, userID string) (map[string]interface{}, common.ErrorCode, error) { + + // 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 + } + + result := map[string]interface{}{ + "id": instance.ID, + "instanceName": instance.InstanceName, + "providerID": instance.ProviderID, + "status": instance.Status, + } + + return result, common.CodeSuccess, nil +} + +func (m *ModelProviderService) AlterProviderInstance(providerName, instanceName, newInstanceName, apiKey, userID string) (common.ErrorCode, error) { + return common.CodeSuccess, nil +} +func (m *ModelProviderService) DropProviderInstance(providerName, instanceName, 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, providerName) + if err != nil { + return common.CodeServerError, err + } + + count, err := m.modelInstanceDAO.DeleteByProviderIDAndInstanceName(provider.ID, instanceName) + if err != nil { + return common.CodeServerError, err + } + + if count == 0 { + return common.CodeNotFound, errors.New("provider instance not found") + } + + return common.CodeSuccess, nil +} + +func (m *ModelProviderService) ListInstanceModels(providerName, instanceName, userID string) ([]map[string]interface{}, error) { + // Get tenant ID from user + tenants, err := m.userTenantDAO.GetByUserIDAndRole(userID, "owner") + if err != nil { + return nil, err + } + + if len(tenants) == 0 { + return nil, 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, err + } + + // Get instance + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + if err != nil { + return nil, err + } + + // Get all models for this instance + disabledModels, err := m.modelDAO.GetModelsByInstanceID(instance.ID) + if err != nil { + return nil, err + } + + // insert models name into a set + modelNames := make(map[string]bool) + for _, model := range disabledModels { + 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" + } else { + model["status"] = "enabled" + } + + } + + return allModels, nil +} + +func (m *ModelProviderService) UpdateModelStatus(providerName, instanceName, modelName, userID, status 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, providerName) + if err != nil { + return common.CodeServerError, err + } + + instance, err := m.modelInstanceDAO.GetByProviderIDAndInstanceName(provider.ID, instanceName) + if err != nil { + return common.CodeServerError, err + } + + model, err := m.modelDAO.GetModelByProviderIDAndInstanceIDAndModelName(provider.ID, instance.ID, modelName) + if err != nil { + var modelID string + modelID, err = generateUUID1Hex() + if err != nil { + return common.CodeServerError, errors.New("fail to get UUID") + } + // Get model info from provider + model = &entity.TenantModel{ + ID: modelID, + ModelName: modelName, + ModelType: model.ModelType, + ProviderID: provider.ID, + InstanceID: instance.ID, + Status: status, + } + err = m.modelDAO.Create(model) + if err != nil { + return common.CodeServerError, errors.New("fail to create model") + } + return common.CodeSuccess, nil + } + + count, err := m.modelDAO.DeleteByModelID(model.ID) + if err != nil { + return common.CodeServerError, err + } + if count == 0 { + return common.CodeNotFound, errors.New("model not found") + } + + return common.CodeSuccess, nil +} + +func (m *ModelProviderService) ChatToModel(providerName, instanceName, modelName, userID, message string) (*string, common.ErrorCode, error) { + + // 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 + } + + _, 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") + } + + _, 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)) + } + + var response string + response, err = providerInfo.ModelDriver.Chat(&modelName, &instance.APIKey, &message, nil) + if err != nil { + return nil, common.CodeServerError, err + } + + return &response, common.CodeSuccess, nil + } + + return nil, common.CodeServerError, errors.New("model is disabled") +}