diff --git a/conf/llm_factories.json b/conf/llm_factories.json index 7095e08ed2..4682bb6efd 100644 --- a/conf/llm_factories.json +++ b/conf/llm_factories.json @@ -1797,6 +1797,14 @@ "llm": [], "rank": "985" }, + { + "name": "MWS", + "logo": "", + "tags": "LLM,TEXT EMBEDDING,TEXT RE-RANK", + "status": "1", + "llm": [], + "rank": "984" + }, { "name": "VLLM", "logo": "", diff --git a/conf/models/mws.json b/conf/models/mws.json new file mode 100644 index 0000000000..e099379535 --- /dev/null +++ b/conf/models/mws.json @@ -0,0 +1,15 @@ +{ + "name": "MWS", + "rank": 984, + "url": { + "default": "" + }, + "url_suffix": { + "models": "openai/v1/models", + "chat": "openai/v1/chat/completions", + "embedding": "openai/v1/embeddings", + "rerank": "cohere/v2/rerank" + }, + "class": "mws", + "models": [] +} diff --git a/docs/guides/models/supported_models.mdx b/docs/guides/models/supported_models.mdx index 16a11e4408..de274f119b 100644 --- a/docs/guides/models/supported_models.mdx +++ b/docs/guides/models/supported_models.mdx @@ -48,6 +48,7 @@ A complete list of model providers supported by RAGFlow, which will continue to | Mistral | `https://mistral.ai` | | ModelScope | `https://www.modelscope.cn` | | Moonshot | `https://www.moonshot.cn` | +| MWS | `https://mws.ru/docs/cloud-platform/gpt.html` | | NovitaAI | `https://novita.ai` | | NVIDIA | `https://www.nvidia.com` | | Ollama | `https://ollama.com` | diff --git a/internal/entity/models/factory.go b/internal/entity/models/factory.go index 1aa4d1da60..87ab9a1ef2 100644 --- a/internal/entity/models/factory.go +++ b/internal/entity/models/factory.go @@ -57,6 +57,8 @@ func (f *ModelFactory) CreateModelDriver(providerName string, baseURL map[string return NewVllmModel(baseURL, urlSuffix), nil case "openai-api-compatible": return NewOpenAIAPICompatibleModel(baseURL, urlSuffix), nil + case "mws": + return NewMWSModel(baseURL, urlSuffix), nil case "xai": return NewXAIModel(baseURL, urlSuffix), nil case "lm-studio": diff --git a/internal/entity/models/mws.go b/internal/entity/models/mws.go new file mode 100644 index 0000000000..75db29b6dc --- /dev/null +++ b/internal/entity/models/mws.go @@ -0,0 +1,316 @@ +// +// 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 ( + "context" + "encoding/json" + "fmt" + "io" + "net/url" + "strings" + + "ragflow/internal/common" +) + +// MWSModel implements the MWS GPT Model Hub chat, embedding, and reranking APIs. +type MWSModel struct { + *DummyModel +} + +// NewMWSModel creates an MWS model driver. +func NewMWSModel(baseURL map[string]string, urlSuffix URLSuffix) *MWSModel { + driver := NewDummyModel(baseURL, urlSuffix) + driver.baseModel.httpClient = NewDriverHTTPClient(true) + return &MWSModel{DummyModel: driver} +} + +// Name returns the public provider identifier. +func (m *MWSModel) Name() string { + return "MWS" +} + +// NewInstance creates an MWS driver with tenant-specific base URLs. +func (m *MWSModel) NewInstance(baseURL map[string]string) ModelDriver { + return NewMWSModel(baseURL, m.baseModel.URLSuffix) +} + +func normalizeMWSProjectURL(rawURL string) (string, error) { + value := strings.TrimSuffix(strings.TrimSpace(rawURL), "/") + parsed, err := url.Parse(value) + if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Hostname() == "" { + return "", fmt.Errorf("MWS API URL must be a project root in the form https://gpt.mwsapis.ru/projects/") + } + parts := strings.Split(strings.Trim(parsed.Path, "/"), "/") + if len(parts) != 2 || parts[0] != "projects" || parts[1] == "" { + return "", fmt.Errorf("MWS API URL must be a project root in the form https://gpt.mwsapis.ru/projects/") + } + if parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" || parsed.User != nil { + return "", fmt.Errorf("MWS API URL must not contain credentials, a query string, or a fragment") + } + parsed.Path = "/projects/" + parts[1] + parsed.RawPath = "" + return parsed.String(), nil +} + +func buildMWSEndpoint(projectURL, endpoint string) (string, error) { + root, err := normalizeMWSProjectURL(projectURL) + if err != nil { + return "", err + } + return root + "/" + strings.Trim(endpoint, "/"), nil +} + +func (m *MWSModel) endpoint(apiConfig *APIConfig, endpoint string) (string, error) { + if err := m.baseModel.APIConfigCheck(apiConfig); err != nil { + return "", err + } + baseURL, err := m.baseModel.GetBaseURL(apiConfig) + if err != nil { + return "", err + } + return buildMWSEndpoint(baseURL, endpoint) +} + +func (m *MWSModel) ListModels(ctx context.Context, apiConfig *APIConfig) ([]ListModelResponse, error) { + endpoint, err := m.endpoint(apiConfig, "openai/v1/models") + if err != nil { + return nil, err + } + body, err := m.baseModel.doGetRequest(ctx, endpoint, apiConfig, nonStreamCallTimeout) + if err != nil { + return nil, err + } + + var modelList ModelList + if err = json.Unmarshal(body, &modelList); err != nil { + return nil, fmt.Errorf("failed to parse MWS model list: %w", err) + } + if modelList.Models == nil { + return nil, fmt.Errorf("invalid MWS models list format") + } + + models := make([]ListModelResponse, 0, len(modelList.Models)) + for _, model := range modelList.Models { + modelTypes := InferModelTypes(model.ID) + if len(modelTypes) != 1 || (modelTypes[0] != "chat" && modelTypes[0] != "embedding" && modelTypes[0] != "rerank") { + continue + } + models = append(models, ListModelResponse{Name: model.ID, ModelTypes: modelTypes}) + } + return models, nil +} + +func buildMWSChatMessages(messages []Message) ([]map[string]any, error) { + if len(messages) == 0 { + return nil, fmt.Errorf("messages are required") + } + result := make([]map[string]any, 0, len(messages)) + for _, message := range messages { + if message.Role != "system" && message.Role != "user" && message.Role != "assistant" { + return nil, fmt.Errorf("unsupported MWS chat message role %q", message.Role) + } + content, ok := message.Content.(string) + if !ok { + return nil, fmt.Errorf("MWS chat message content must be a string") + } + result = append(result, map[string]any{ + "role": message.Role, + "content": content, + }) + } + return result, nil +} + +func buildMWSChatRequest(modelName string, messages []Message, config *ChatConfig, stream bool) (map[string]any, error) { + modelName = strings.TrimSpace(modelName) + if modelName == "" { + return nil, fmt.Errorf("model name is required") + } + chatMessages, err := buildMWSChatMessages(messages) + if err != nil { + return nil, err + } + request := map[string]any{ + "model": modelName, + "messages": chatMessages, + } + if config != nil { + if config.Temperature != nil { + request["temperature"] = *config.Temperature + } + if config.MaxTokens != nil { + request["max_completion_tokens"] = *config.MaxTokens + } + } + if stream { + request["stream"] = true + request["stream_options"] = map[string]any{"include_usage": true} + } + return request, nil +} + +// ChatWithMessages sends a non-streaming MWS Chat Completions request. +func (m *MWSModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) { + endpoint, err := m.endpoint(apiConfig, "openai/v1/chat/completions") + if err != nil { + return nil, err + } + request, err := buildMWSChatRequest(modelName, messages, chatConfig, false) + if err != nil { + return nil, err + } + body, err := m.baseModel.doRequest(ctx, endpoint, apiConfig, request, nonStreamCallTimeout) + if err != nil { + return nil, err + } + return HandleNonStreamingResponse(body, modelUsage, chatConfig, OpenAIParserConfig) +} + +// ChatStreamlyWithSender sends a streaming MWS Chat Completions request. +func (m *MWSModel) ChatStreamlyWithSender(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatConfig *ChatConfig, modelUsage *common.ModelUsage, sender func(*string, *string) error) error { + if sender == nil { + return fmt.Errorf("sender is required") + } + if err := validateStreamConfig(chatConfig); err != nil { + return err + } + endpoint, err := m.endpoint(apiConfig, "openai/v1/chat/completions") + if err != nil { + return err + } + request, err := buildMWSChatRequest(modelName, messages, chatConfig, true) + if err != nil { + return err + } + return m.baseModel.doStreamRequest(ctx, endpoint, apiConfig, request, streamCallTimeout, func(body io.ReadCloser) error { + return HandleStreamingResponse(body, modelUsage, chatConfig, OpenAIParserConfig, sender) + }) +} + +type mwsEmbeddingResponse struct { + Data []struct { + Index int `json:"index"` + Embedding []float64 `json:"embedding"` + } `json:"data"` + Usage TokenUsage `json:"usage"` +} + +func (m *MWSModel) Embed(ctx context.Context, modelName *string, texts []string, apiConfig *APIConfig, _ *EmbeddingConfig, modelUsage *common.ModelUsage) ([]EmbeddingData, error) { + endpoint, err := m.endpoint(apiConfig, "openai/v1/embeddings") + if err != nil { + return nil, err + } + if len(texts) == 0 { + return []EmbeddingData{}, nil + } + if modelName == nil || strings.TrimSpace(*modelName) == "" { + return nil, fmt.Errorf("model name is required") + } + + body, err := m.baseModel.doRequest(ctx, endpoint, apiConfig, map[string]any{ + "model": *modelName, + "input": texts, + }, nonStreamCallTimeout) + if err != nil { + return nil, err + } + + var response mwsEmbeddingResponse + if err = json.Unmarshal(body, &response); err != nil { + return nil, fmt.Errorf("failed to parse MWS embedding response: %w", err) + } + if len(response.Data) != len(texts) { + return nil, fmt.Errorf("MWS embedding response returned %d vectors for %d inputs", len(response.Data), len(texts)) + } + + embeddings := make([]EmbeddingData, len(texts)) + seen := make([]bool, len(texts)) + for _, item := range response.Data { + if item.Index < 0 || item.Index >= len(texts) || seen[item.Index] { + return nil, fmt.Errorf("unexpected MWS embedding index %d for %d inputs", item.Index, len(texts)) + } + seen[item.Index] = true + embeddings[item.Index] = EmbeddingData{Index: item.Index, Embedding: item.Embedding} + } + recordResponseUsage(modelUsage, "", &response.Usage, "embedding") + return embeddings, nil +} + +type mwsRerankResponse struct { + ID string `json:"id"` + Results []struct { + Index int `json:"index"` + RelevanceScore float64 `json:"relevance_score"` + } `json:"results"` + Meta struct { + Tokens struct { + InputTokens int `json:"input_tokens"` + } `json:"tokens"` + } `json:"meta"` +} + +func (m *MWSModel) Rerank(ctx context.Context, modelName *string, query string, documents []string, apiConfig *APIConfig, _ *RerankConfig, modelUsage *common.ModelUsage) (*RerankResponse, error) { + endpoint, err := m.endpoint(apiConfig, "cohere/v2/rerank") + if err != nil { + return nil, err + } + if len(documents) == 0 { + return &RerankResponse{}, nil + } + if modelName == nil || strings.TrimSpace(*modelName) == "" { + return nil, fmt.Errorf("model name is required") + } + + body, err := m.baseModel.doRequest(ctx, endpoint, apiConfig, map[string]any{ + "model": *modelName, + "query": query, + "documents": documents, + "top_n": len(documents), + }, nonStreamCallTimeout) + if err != nil { + return nil, err + } + + var response mwsRerankResponse + if err = json.Unmarshal(body, &response); err != nil { + return nil, fmt.Errorf("failed to parse MWS rerank response: %w", err) + } + if len(response.Results) != len(documents) { + return nil, fmt.Errorf("MWS rerank response returned %d results for %d documents", len(response.Results), len(documents)) + } + ranked := make([]RerankResult, len(documents)) + seen := make([]bool, len(documents)) + for i := range ranked { + ranked[i].Index = i + } + for _, item := range response.Results { + if item.Index < 0 || item.Index >= len(documents) || seen[item.Index] { + return nil, fmt.Errorf("unexpected MWS rerank index %d for %d documents", item.Index, len(documents)) + } + seen[item.Index] = true + ranked[item.Index].RelevanceScore = item.RelevanceScore + } + usage := TokenUsage{PromptTokens: response.Meta.Tokens.InputTokens, TotalTokens: response.Meta.Tokens.InputTokens} + recordResponseUsage(modelUsage, response.ID, &usage, "rerank") + return &RerankResponse{Data: ranked}, nil +} + +func (m *MWSModel) CheckConnection(ctx context.Context, apiConfig *APIConfig) error { + _, err := m.ListModels(ctx, apiConfig) + return err +} diff --git a/internal/entity/models/mws_test.go b/internal/entity/models/mws_test.go new file mode 100644 index 0000000000..2419bacd9c --- /dev/null +++ b/internal/entity/models/mws_test.go @@ -0,0 +1,326 @@ +// +// 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 ( + "context" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "sync/atomic" + "testing" +) + +func newMWSTestDriver(serverURL string) (*MWSModel, *APIConfig) { + token := "token" + baseURL := serverURL + "/projects/test-project" + driver := NewMWSModel(map[string]string{"default": baseURL}, URLSuffix{}) + return driver, &APIConfig{ApiKey: &token, BaseURL: &baseURL} +} + +func decodeMWSRequest(t *testing.T, request *http.Request) map[string]any { + t.Helper() + if request.Method != http.MethodPost { + t.Fatalf("unexpected method: %s", request.Method) + } + if request.Header.Get("Authorization") != "Bearer token" { + t.Fatalf("unexpected authorization: %q", request.Header.Get("Authorization")) + } + if request.Header.Get("Content-Type") != "application/json" { + t.Fatalf("unexpected content type: %q", request.Header.Get("Content-Type")) + } + body, err := io.ReadAll(request.Body) + if err != nil { + t.Fatalf("read request body: %v", err) + } + var payload map[string]any + if err = json.Unmarshal(body, &payload); err != nil { + t.Fatalf("decode request body: %v", err) + } + return payload +} + +func TestNormalizeMWSProjectURL(t *testing.T) { + got, err := normalizeMWSProjectURL("https://gpt.mwsapis.ru/projects/demo/") + if err != nil { + t.Fatalf("normalize URL: %v", err) + } + if got != "https://gpt.mwsapis.ru/projects/demo" { + t.Fatalf("unexpected normalized URL: %s", got) + } + if _, err = normalizeMWSProjectURL("https://gpt.mwsapis.ru/projects/demo/openai/v1"); err == nil { + t.Fatal("expected a non-root URL to be rejected") + } +} + +func TestMWSListModelsUsesOpenAIEndpointAndFiltersTypes(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.Method != http.MethodGet || request.URL.Path != "/projects/test-project/openai/v1/models" { + t.Fatalf("unexpected request: %s %s", request.Method, request.URL.Path) + } + if request.Header.Get("Authorization") != "Bearer token" { + t.Fatalf("unexpected authorization: %q", request.Header.Get("Authorization")) + } + body, _ := io.ReadAll(request.Body) + if len(body) != 0 { + t.Fatalf("GET models request must not have a body: %q", body) + } + response.Header().Set("Content-Type", "application/json") + _, _ = response.Write([]byte(`{"object":"list","data":[{"id":"bge-m3"},{"id":"bge-reranker-v2-m3"},{"id":"qwen3-32b"},{"id":"qwen-vl"}]}`)) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + models, err := driver.ListModels(context.Background(), apiConfig) + if err != nil { + t.Fatalf("list models: %v", err) + } + want := []ListModelResponse{ + {Name: "bge-m3", ModelTypes: []string{"embedding"}}, + {Name: "bge-reranker-v2-m3", ModelTypes: []string{"rerank"}}, + {Name: "qwen3-32b", ModelTypes: []string{"chat"}}, + } + if !reflect.DeepEqual(models, want) { + t.Fatalf("unexpected models: %#v", models) + } +} + +func TestMWSChatSendsOnlyDocumentedFields(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/projects/test-project/openai/v1/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + payload := decodeMWSRequest(t, request) + want := map[string]any{ + "model": "qwen3-32b", + "messages": []any{ + map[string]any{"role": "system", "content": "Be concise."}, + map[string]any{"role": "user", "content": "Hello"}, + }, + "temperature": 0.25, + "max_completion_tokens": float64(128), + } + if !reflect.DeepEqual(payload, want) { + t.Fatalf("unexpected chat payload: %#v", payload) + } + response.Header().Set("Content-Type", "application/json") + _, _ = response.Write([]byte(`{"id":"chat-1","model":"qwen3-32b","choices":[{"index":0,"finish_reason":"stop","message":{"role":"assistant","content":"Hi"}}],"usage":{"prompt_tokens":4,"completion_tokens":1,"total_tokens":5}}`)) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + temperature := 0.25 + maxTokens := 128 + topP := 0.9 + stop := []string{"ignored"} + response, err := driver.ChatWithMessages( + context.Background(), + "qwen3-32b", + []Message{ + {Role: "system", Content: "Be concise."}, + {Role: "user", Content: "Hello", ToolCallID: "ignored"}, + }, + apiConfig, + &ChatConfig{Temperature: &temperature, MaxTokens: &maxTokens, TopP: &topP, Stop: &stop, Tools: map[string]any{"ignored": true}}, + nil, + ) + if err != nil { + t.Fatalf("chat: %v", err) + } + if response.Answer == nil || *response.Answer != "Hi" || response.Usage == nil || response.Usage.TotalTokens != 5 { + t.Fatalf("unexpected chat response: %#v", response) + } +} + +func TestMWSChatStreamingUsesDocumentedFields(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/projects/test-project/openai/v1/chat/completions" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + payload := decodeMWSRequest(t, request) + want := map[string]any{ + "model": "qwen3-32b", + "messages": []any{map[string]any{"role": "user", "content": "Hello"}}, + "stream": true, + "stream_options": map[string]any{ + "include_usage": true, + }, + } + if !reflect.DeepEqual(payload, want) { + t.Fatalf("unexpected streaming chat payload: %#v", payload) + } + response.Header().Set("Content-Type", "text/event-stream") + _, _ = io.WriteString(response, "data: {\"model\":\"qwen3-32b\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hel\"}}]}\n\n") + _, _ = io.WriteString(response, "data: {\"model\":\"qwen3-32b\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"lo\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":4,\"completion_tokens\":1,\"total_tokens\":5}}\n\n") + _, _ = io.WriteString(response, "data: [DONE]\n\n") + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + stream := true + config := &ChatConfig{Stream: &stream} + var chunks []string + err := driver.ChatStreamlyWithSender( + context.Background(), + "qwen3-32b", + []Message{{Role: "user", Content: "Hello"}}, + apiConfig, + config, + nil, + func(content, _ *string) error { + if content != nil { + chunks = append(chunks, *content) + } + return nil + }, + ) + if err != nil { + t.Fatalf("stream chat: %v", err) + } + if !reflect.DeepEqual(chunks, []string{"Hel", "lo", "[DONE]"}) { + t.Fatalf("unexpected stream chunks: %#v", chunks) + } + if config.UsageResult == nil || config.UsageResult.TotalTokens != 5 { + t.Fatalf("unexpected stream usage: %#v", config.UsageResult) + } +} + +func TestMWSEmbedSendsOnlyDocumentedFieldsAndOrdersVectors(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/projects/test-project/openai/v1/embeddings" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + payload := decodeMWSRequest(t, request) + want := map[string]any{"model": "bge-m3", "input": []any{"first", "second"}} + if !reflect.DeepEqual(payload, want) { + t.Fatalf("unexpected embedding payload: %#v", payload) + } + response.Header().Set("Content-Type", "application/json") + _, _ = response.Write([]byte(`{"data":[{"index":1,"embedding":[0.3,0.4]},{"index":0,"embedding":[0.1,0.2]}],"usage":{"prompt_tokens":7,"total_tokens":7}}`)) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + modelName := "bge-m3" + embeddings, err := driver.Embed(context.Background(), &modelName, []string{"first", "second"}, apiConfig, &EmbeddingConfig{Dimension: 256}, nil) + if err != nil { + t.Fatalf("embed: %v", err) + } + if len(embeddings) != 2 || embeddings[0].Index != 0 || embeddings[1].Index != 1 || embeddings[0].Embedding[0] != 0.1 || embeddings[1].Embedding[0] != 0.3 { + t.Fatalf("unexpected embeddings: %#v", embeddings) + } +} + +func TestMWSRerankUsesCohereEndpointAndOriginalIndexOrder(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.URL.Path != "/projects/test-project/cohere/v2/rerank" { + t.Fatalf("unexpected path: %s", request.URL.Path) + } + payload := decodeMWSRequest(t, request) + want := map[string]any{ + "model": "bge-reranker-v2-m3", + "query": "query", + "documents": []any{"first", "second"}, + "top_n": float64(2), + } + if !reflect.DeepEqual(payload, want) { + t.Fatalf("unexpected rerank payload: %#v", payload) + } + response.Header().Set("Content-Type", "application/json") + _, _ = response.Write([]byte(`{"id":"score-1","results":[{"index":1,"relevance_score":0.9},{"index":0,"relevance_score":0.2}],"meta":{"tokens":{"input_tokens":9}}}`)) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + modelName := "bge-reranker-v2-m3" + result, err := driver.Rerank(context.Background(), &modelName, "query", []string{"first", "second"}, apiConfig, &RerankConfig{TopN: 1}, nil) + if err != nil { + t.Fatalf("rerank: %v", err) + } + if len(result.Data) != 2 || result.Data[0].Index != 0 || result.Data[0].RelevanceScore != 0.2 || result.Data[1].Index != 1 || result.Data[1].RelevanceScore != 0.9 { + t.Fatalf("unexpected rerank result: %#v", result.Data) + } +} + +func TestMWSRejectsEmptyTokenWithoutRequest(t *testing.T) { + var calls atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + calls.Add(1) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + empty := " " + apiConfig.ApiKey = &empty + modelName := "bge-m3" + _, err := driver.Embed(context.Background(), &modelName, []string{"text"}, apiConfig, nil, nil) + if err == nil || !strings.Contains(err.Error(), "api key is required") { + t.Fatalf("expected an API key error, got %v", err) + } + if calls.Load() != 0 { + t.Fatalf("unexpected HTTP calls: %d", calls.Load()) + } +} + +func TestMWSEmptyInputsDoNotSendRequests(t *testing.T) { + var calls atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + calls.Add(1) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + modelName := "bge-m3" + embeddings, err := driver.Embed(context.Background(), &modelName, []string{}, apiConfig, nil, nil) + if err != nil || len(embeddings) != 0 { + t.Fatalf("empty embedding input: data=%#v err=%v", embeddings, err) + } + ranked, err := driver.Rerank(context.Background(), &modelName, "query", []string{}, apiConfig, nil, nil) + if err != nil || len(ranked.Data) != 0 { + t.Fatalf("empty rerank input: data=%#v err=%v", ranked, err) + } + if calls.Load() != 0 { + t.Fatalf("unexpected HTTP calls: %d", calls.Load()) + } +} + +func TestMWSErrorResponseIsReturned(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) { + http.Error(response, "MWS unavailable", http.StatusServiceUnavailable) + })) + defer server.Close() + + driver, apiConfig := newMWSTestDriver(server.URL) + modelName := "bge-m3" + _, err := driver.Embed(context.Background(), &modelName, []string{"text"}, apiConfig, nil, nil) + if err == nil || !strings.Contains(err.Error(), "status 503") || !strings.Contains(err.Error(), "MWS unavailable") { + t.Fatalf("unexpected MWS error: %v", err) + } +} + +func TestMWSFactoryRegistration(t *testing.T) { + driver, err := NewModelFactory().CreateModelDriver("MWS", map[string]string{"default": "https://gpt.mwsapis.ru/projects/demo"}, URLSuffix{}) + if err != nil { + t.Fatalf("create driver: %v", err) + } + if driver.Name() != "MWS" { + t.Fatalf("unexpected driver: %s", driver.Name()) + } +} diff --git a/rag/llm/chat_model.py b/rag/llm/chat_model.py index 5393b449d0..bb50d49c1f 100644 --- a/rag/llm/chat_model.py +++ b/rag/llm/chat_model.py @@ -24,6 +24,7 @@ from abc import ABC from copy import deepcopy from urllib.parse import urljoin +import aiohttp import json_repair from json.decoder import JSONDecodeError import litellm @@ -36,6 +37,7 @@ from common.llm_request_context import current_llm_user from common.token_utils import num_tokens_from_string, total_token_count_from_response, usage_from_response from rag.llm import FACTORY_DEFAULT_BASE_URL, LITELLM_PROVIDER_PREFIX, SupportedLiteLLMProvider from rag.llm.key_utils import _normalize_replicate_key +from rag.llm.mws_utils import mws_api_url, require_mws_token from rag.llm.tool_decorator import FunctionToolSession, is_tool from rag.nlp import is_chinese, is_english from rag.utils.url_utils import ensure_v1 @@ -1069,6 +1071,137 @@ class OpenAI_APIChat(Base): super().__init__(key, model_name, base_url, **kwargs) +class MWSChat(Base): + """MWS Chat Completions adapter with a documentation-only request body.""" + + _FACTORY_NAME = "MWS" + _ROLES = {"system", "user", "assistant"} + + def __init__(self, key, model_name, base_url, **kwargs): + """Initialize chat access for an MWS project and model deployment.""" + token = require_mws_token(key) + self.chat_url = mws_api_url(base_url, "openai/v1/chat/completions") + self.headers = { + "Content-Type": "application/json", + "Authorization": f"Bearer {token}", + } + super().__init__( + token, + model_name.split("___")[0], + mws_api_url(base_url, "openai/v1"), + **kwargs, + ) + + def _clean_conf(self, gen_conf): + """Keep only generation parameters documented by the MWS API.""" + gen_conf = gen_conf or {} + cleaned = {} + if gen_conf.get("temperature") is not None: + cleaned["temperature"] = gen_conf["temperature"] + max_tokens = gen_conf.get("max_completion_tokens") + if max_tokens is None: + max_tokens = gen_conf.get("max_tokens") + if max_tokens is not None: + cleaned["max_completion_tokens"] = max_tokens + return cleaned + + def _request_body(self, history, gen_conf, *, stream): + """Build a strict MWS chat request from RAGFlow messages and options.""" + messages = [] + for message in history: + role = message.get("role") if isinstance(message, dict) else None + content = message.get("content") if isinstance(message, dict) else None + if role not in self._ROLES or not isinstance(content, str): + raise ValueError("MWS chat messages must contain only a system, user, or assistant role and string content") + messages.append({"role": role, "content": content}) + if not messages: + raise ValueError("MWS chat messages are required") + + body = {"model": self.model_name, "messages": messages} + body.update(self._clean_conf(gen_conf)) + if stream: + body["stream"] = True + body["stream_options"] = {"include_usage": True} + return body + + async def _post_json(self, body): + """Send a non-streaming MWS chat request and decode its JSON response.""" + timeout = aiohttp.ClientTimeout(total=int(os.environ.get("LLM_TIMEOUT_SECONDS", 600))) + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.post( + self.chat_url, + headers=self.headers, + json=body, + ) as response: + if response.status != 200: + raise RuntimeError(f"MWS chat request failed with status {response.status}: {await response.text()}") + return await response.json() + + async def _async_chat(self, history, gen_conf, **kwargs): + """Return one complete MWS chat answer together with its token usage.""" + payload = await self._post_json(self._request_body(history, gen_conf, stream=False)) + self.last_usage = usage_from_response(payload) + choices = payload.get("choices") if isinstance(payload, dict) else None + if not isinstance(choices, list) or not choices: + raise ValueError("MWS chat response does not contain choices") + choice = choices[0] + message = choice.get("message") if isinstance(choice, dict) else None + content = message.get("content") if isinstance(message, dict) else None + if not isinstance(content, str): + raise ValueError("MWS chat response does not contain message content") + answer = content.strip() + if choice.get("finish_reason") == "length": + answer = self._length_stop(answer) + return answer, total_token_count_from_response(payload) + + async def _async_chat_streamly(self, history, gen_conf, **kwargs): + """Yield MWS SSE content chunks and attach usage to the final chunk.""" + body = self._request_body(history, gen_conf, stream=True) + timeout = aiohttp.ClientTimeout(total=int(os.environ.get("LLM_TIMEOUT_SECONDS", 600))) + pending_content = None + estimated_tokens = 0 + reported_tokens = 0 + + async with aiohttp.ClientSession(timeout=timeout) as session: + async with session.post( + self.chat_url, + headers=self.headers, + json=body, + ) as response: + if response.status != 200: + raise RuntimeError(f"MWS chat request failed with status {response.status}: {await response.text()}") + + async for raw_line in response.content: + line = raw_line.decode("utf-8").strip() + if not line.startswith("data:"): + continue + data = line[5:].strip() + if data == "[DONE]": + break + event = json.loads(data) + usage = usage_from_response(event) + if usage["total_tokens"]: + self.last_usage = usage + reported_tokens = usage["total_tokens"] + + choices = event.get("choices") + if not isinstance(choices, list) or not choices: + continue + choice = choices[0] + delta = choice.get("delta") if isinstance(choice, dict) else None + content = delta.get("content") if isinstance(delta, dict) else None + if not isinstance(content, str) or not content: + continue + if choice.get("finish_reason") == "length": + content = self._length_stop(content) + if pending_content is not None: + yield pending_content, 0 + pending_content = content + estimated_tokens += num_tokens_from_string(content) + + yield pending_content or "", reported_tokens or estimated_tokens + + class Xiaomi(Base): _FACTORY_NAME = "Xiaomi" diff --git a/rag/llm/embedding_model.py b/rag/llm/embedding_model.py index e4790e0222..5da4195814 100644 --- a/rag/llm/embedding_model.py +++ b/rag/llm/embedding_model.py @@ -33,6 +33,7 @@ from common.aimlapi_utils import attribution_headers from common.exceptions import ModelException from common.token_utils import num_tokens_from_string, truncate, total_token_count_from_response from rag.llm.key_utils import _normalize_replicate_key +from rag.llm.mws_utils import mws_api_url, require_mws_token from rag.utils.url_utils import append_api_path, ensure_v1 import logging import base64 @@ -880,6 +881,41 @@ class OpenAI_APIEmbed(OpenAIEmbed): self.model_name = model_name.split("___")[0] +class MWSEmbed(OpenAIEmbed): + """MWS embedding adapter with the exact documented request body.""" + + _FACTORY_NAME = "MWS" + + def __init__(self, key, model_name, base_url): + """Initialize embedding access for an MWS project and deployment.""" + self.api_key = require_mws_token(key) + self.base_url = mws_api_url(base_url, "openai/v1/embeddings") + self.model_name = model_name.split("___")[0] + + def _call(self, batch): + """Embed a batch and restore vectors to their original input order.""" + response = requests.post( + self.base_url, + headers={"Content-Type": "application/json", "Authorization": f"Bearer {self.api_key}"}, + json={"model": self.model_name, "input": batch}, + timeout=30, + ) + _raise_model_exception_if_failed(response) + payload = response.json() + data = payload.get("data") if isinstance(payload, dict) else None + if not isinstance(data, list) or len(data) != len(batch): + count = len(data) if isinstance(data, list) else 0 + raise ValueError(f"MWS returned {count} embeddings for {len(batch)} inputs") + + embeddings = [None] * len(batch) + for item in data: + index = item.get("index") if isinstance(item, dict) else None + if not isinstance(index, int) or isinstance(index, bool) or index < 0 or index >= len(batch) or embeddings[index] is not None: + raise ValueError(f"unexpected MWS embedding index: {index}") + embeddings[index] = item["embedding"] + return embeddings, total_token_count_from_response(payload) + + class GreenPTEmbed(OpenAIEmbed): """GreenPT OpenAI-compatible embedding adapter.""" diff --git a/rag/llm/model_meta.py b/rag/llm/model_meta.py index ae2dbde7b1..51834f5787 100644 --- a/rag/llm/model_meta.py +++ b/rag/llm/model_meta.py @@ -24,6 +24,7 @@ from typing import ClassVar from common.aimlapi_utils import attribution_headers from common.constants import LLMType +from rag.llm.mws_utils import mws_api_url, normalize_mws_project_url, require_mws_token class Base(ABC): @@ -462,6 +463,131 @@ class OpenAIAPICompatible(Base): return model_list +class MWS(OpenAIAPICompatible): + """Discover supported MWS deployments through the project models API.""" + + _FACTORY_NAME = "MWS" + + def __init__(self, api_key: str, base_url: str = None): + """Initialize dynamic model discovery for an MWS project.""" + try: + token = require_mws_token(api_key) + except ValueError as error: + logging.warning( + "mws_model_discovery_validation_failed", + extra={ + "provider": self._FACTORY_NAME, + "operation": "model_discovery", + "validation_target": "token", + "error_type": type(error).__name__, + }, + ) + raise + + try: + project_url = normalize_mws_project_url(base_url) + except ValueError as error: + logging.warning( + "mws_model_discovery_validation_failed", + extra={ + "provider": self._FACTORY_NAME, + "operation": "model_discovery", + "validation_target": "api_url", + "error_type": type(error).__name__, + }, + ) + raise + + super().__init__(token, project_url) + + def _get_model_list_url(self): + """Return the OpenAI-compatible models endpoint for the MWS project.""" + return mws_api_url(self.base_url, "openai/v1/models") + + def _format_model_list(self, raw_model_list): + """Keep only chat, embedding, and reranking MWS deployments.""" + supported_types = { + LLMType.CHAT.value, + LLMType.EMBEDDING.value, + LLMType.RERANK.value, + } + return [model for model in super()._format_model_list(raw_model_list) if len(model.get("model_types") or []) == 1 and model["model_types"][0] in supported_types] + + async def get_model_list(self): + """Discover MWS models while logging safe request and result metadata.""" + try: + url = self._get_model_list_url() + except ValueError as error: + logging.warning( + "mws_model_discovery_validation_failed", + extra={ + "provider": self._FACTORY_NAME, + "operation": "model_discovery", + "validation_target": "api_url", + "error_type": type(error).__name__, + }, + ) + raise + + log_context = { + "provider": self._FACTORY_NAME, + "operation": "model_discovery", + "url": url, + } + logging.info("mws_model_discovery_request", extra=log_context) + + try: + async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30)) as session: + async with session.get(url, headers={"Authorization": f"Bearer {self._get_api_key()}"}) as response: + if response.status != 200: + logging.warning( + "mws_model_discovery_request_failed", + extra={ + **log_context, + "failure_stage": "http_response", + "http_status": response.status, + }, + ) + logging.info( + "mws_model_discovery_completed", + extra={**log_context, "result_count": 0}, + ) + return [] + raw_model_list = await response.json() + except Exception as error: + logging.warning( + "mws_model_discovery_request_failed", + extra={ + **log_context, + "failure_stage": "request", + "error_type": type(error).__name__, + }, + ) + raise + + if not raw_model_list: + models = [] + else: + try: + models = self._format_model_list(raw_model_list) + except Exception as error: + logging.warning( + "mws_model_discovery_validation_failed", + extra={ + **log_context, + "validation_target": "response", + "error_type": type(error).__name__, + }, + ) + raise + + logging.info( + "mws_model_discovery_completed", + extra={**log_context, "result_count": len(models)}, + ) + return models + + class NVIDIA(OpenAIAPICompatible): _FACTORY_NAME = "NVIDIA" _HOSTED_API_HOST = "integrate.api.nvidia.com" diff --git a/rag/llm/mws_utils.py b/rag/llm/mws_utils.py new file mode 100644 index 0000000000..2f698b7983 --- /dev/null +++ b/rag/llm/mws_utils.py @@ -0,0 +1,43 @@ +# +# 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. +# +"""Validation and URL helpers for the MWS GPT Model Hub provider.""" + +from urllib.parse import urlparse, urlunparse + + +def normalize_mws_project_url(base_url: str | None) -> str: + """Validate and normalize an MWS GPT Model Hub project-root URL.""" + value = (base_url or "").strip().rstrip("/") + parsed = urlparse(value) + path_parts = parsed.path.strip("/").split("/") + if parsed.scheme not in {"http", "https"} or not parsed.netloc or not parsed.hostname or len(path_parts) != 2 or path_parts[0] != "projects" or not path_parts[1]: + raise ValueError("MWS API URL must be a project root in the form https://gpt.mwsapis.ru/projects/") + if parsed.username or parsed.password or parsed.params or parsed.query or parsed.fragment: + raise ValueError("MWS API URL must not contain credentials, parameters, a query string, or a fragment") + return urlunparse((parsed.scheme, parsed.netloc, f"/projects/{path_parts[1]}", "", "", "")) + + +def mws_api_url(base_url: str | None, endpoint: str) -> str: + """Build an MWS API endpoint relative to a validated project root.""" + return f"{normalize_mws_project_url(base_url)}/{endpoint.strip('/')}" + + +def require_mws_token(token: str | None) -> str: + """Return a normalized MWS bearer token or reject an empty value.""" + value = (token or "").strip() + if not value: + raise ValueError("MWS Token is required") + return value diff --git a/rag/llm/rerank_model.py b/rag/llm/rerank_model.py index 954ff58971..0a9d5eee24 100644 --- a/rag/llm/rerank_model.py +++ b/rag/llm/rerank_model.py @@ -15,6 +15,7 @@ # import json import logging +import math import time from abc import ABC from urllib.parse import urljoin @@ -27,6 +28,7 @@ from yarl import URL from common.log_utils import log_exception from common.token_utils import num_tokens_from_string, truncate, total_token_count_from_response +from rag.llm.mws_utils import mws_api_url, require_mws_token class Base(ABC): @@ -200,6 +202,7 @@ class NvidiaRerank(Base): def __init__(self, key, model_name, base_url="https://ai.api.nvidia.com/v1/retrieval/nvidia/"): if not base_url: base_url = "https://ai.api.nvidia.com/v1/retrieval/nvidia/" + base_url = base_url.rstrip("/") + "/" self.model_name = model_name # Default to NVIDIA's generic reranking endpoint. base_url used to be @@ -210,10 +213,18 @@ class NvidiaRerank(Base): if self.model_name == "nvidia/nv-rerankqa-mistral-4b-v3": self.base_url = urljoin(base_url, "nv-rerankqa-mistral-4b-v3/reranking") - - if self.model_name == "nvidia/rerank-qa-mistral-4b": + elif self.model_name == "nvidia/rerank-qa-mistral-4b": self.base_url = urljoin(base_url, "reranking") self.model_name = "nv-rerank-qa-mistral-4b:1" + else: + logging.info( + "nvidia_rerank_fallback_endpoint_assigned", + extra={ + "provider": self._FACTORY_NAME, + "model": self.model_name, + "endpoint": self.base_url, + }, + ) self.headers = { "accept": "application/json", @@ -287,6 +298,97 @@ class OpenAI_APIRerank(Base): return rank, token_count +class MWSRerank(OpenAI_APIRerank): + """MWS reranker using its Cohere-compatible v2 endpoint.""" + + _FACTORY_NAME = "MWS" + + def __init__(self, key, model_name, base_url): + """Initialize reranking access for an MWS project and deployment.""" + token = require_mws_token(key) + super().__init__(token, model_name, mws_api_url(base_url, "cohere/v2/rerank")) + + def _compute_rank(self, query: str, texts: List) -> Tuple[np.ndarray, int]: + """Score candidate texts and restore scores to document input order.""" + documents = [truncate(text, 500) for text in texts] + log_context = { + "provider": self._FACTORY_NAME, + "operation": "rerank", + "endpoint": self.base_url, + "model": self.model_name, + "document_count": len(documents), + } + logging.info("mws_rerank_request", extra=log_context) + + try: + response = requests.post( + self.base_url, + headers=self.headers, + json={ + "model": self.model_name, + "query": query, + "documents": documents, + "top_n": len(documents), + }, + timeout=30, + ) + response.raise_for_status() + except Exception as error: + logging.warning( + "mws_rerank_failed", + extra={ + **log_context, + "failure_stage": "http_request", + "error_type": type(error).__name__, + }, + ) + raise + + try: + payload = response.json() + except Exception as error: + logging.warning( + "mws_rerank_failed", + extra={ + **log_context, + "failure_stage": "json_parsing", + "error_type": type(error).__name__, + }, + ) + raise + + try: + results = payload.get("results") if isinstance(payload, dict) else None + if not isinstance(results, list) or len(results) != len(documents): + count = len(results) if isinstance(results, list) else 0 + raise ValueError(f"MWS returned {count} rerank results for {len(documents)} documents") + + rank = np.zeros(len(documents), dtype=float) + seen = set() + for item in results: + index = item.get("index") if isinstance(item, dict) else None + if not isinstance(index, int) or isinstance(index, bool) or index < 0 or index >= len(documents) or index in seen: + raise ValueError(f"unexpected MWS rerank index: {index}") + relevance_score = item.get("relevance_score") + if isinstance(relevance_score, bool) or not isinstance(relevance_score, (int, float)) or not math.isfinite(relevance_score): + raise ValueError(f"unexpected MWS rerank relevance_score at index {index}: {relevance_score!r}") + seen.add(index) + rank[index] = relevance_score + except Exception as error: + logging.warning( + "mws_rerank_failed", + extra={ + **log_context, + "failure_stage": "response_validation", + "error_type": type(error).__name__, + }, + ) + raise + + token_count = num_tokens_from_string(query) + sum(num_tokens_from_string(document) for document in documents) + return rank, token_count + + class CoHereRerank(Base): _FACTORY_NAME = ["Cohere", "VLLM"] diff --git a/test/unit_test/rag/llm/test_mws.py b/test/unit_test/rag/llm/test_mws.py new file mode 100644 index 0000000000..d7af6996dd --- /dev/null +++ b/test/unit_test/rag/llm/test_mws.py @@ -0,0 +1,607 @@ +# +# 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. +# +"""Tests for MWS provider registration, discovery, and inference adapters.""" + +from unittest.mock import AsyncMock, MagicMock, call, patch + +import numpy as np +import pytest + +from rag.llm import ChatModel, EmbeddingModel, ModelMeta, RerankModel +from rag.llm.chat_model import MWSChat +from rag.llm.embedding_model import MWSEmbed +from rag.llm.model_meta import MWS +from rag.llm.mws_utils import mws_api_url, normalize_mws_project_url +from rag.llm.rerank_model import MWSRerank + + +PROJECT_URL = "https://gpt.mwsapis.ru/projects/test-project" + + +def _response(payload, status_code=200): + """Create a mocked synchronous HTTP response for MWS adapter tests.""" + response = MagicMock() + response.status_code = status_code + response.json.return_value = payload + response.text = str(payload) + return response + + +def _async_context(value): + """Wrap a mocked value in an asynchronous context manager.""" + context = MagicMock() + context.__aenter__ = AsyncMock(return_value=value) + context.__aexit__ = AsyncMock(return_value=None) + return context + + +@pytest.mark.p1 +def test_mws_provider_registration(): + """Register every supported MWS adapter under the public provider name.""" + assert ChatModel["MWS"] is MWSChat + assert EmbeddingModel["MWS"] is MWSEmbed + assert RerankModel["MWS"] is MWSRerank + assert ModelMeta["MWS"] is MWS + + +@pytest.mark.p1 +def test_mws_project_url_validation_and_endpoints(): + """Normalize project roots and construct each supported MWS endpoint.""" + assert normalize_mws_project_url(PROJECT_URL + "/") == PROJECT_URL + assert mws_api_url(PROJECT_URL, "openai/v1/chat/completions") == PROJECT_URL + "/openai/v1/chat/completions" + assert mws_api_url(PROJECT_URL, "openai/v1/embeddings") == PROJECT_URL + "/openai/v1/embeddings" + assert mws_api_url(PROJECT_URL, "cohere/v2/rerank") == PROJECT_URL + "/cohere/v2/rerank" + with pytest.raises(ValueError, match="project root"): + normalize_mws_project_url("https://gpt.mwsapis.ru/openai/v1") + with pytest.raises(ValueError, match="query string"): + normalize_mws_project_url(PROJECT_URL + "?secret=value") + + +@pytest.mark.p1 +def test_mws_model_discovery_logs_validation_failure_without_credentials(): + """Log the failed validation target without exposing rejected URL data.""" + with patch("rag.llm.model_meta.logging.warning") as warning: + with pytest.raises(ValueError, match="query string"): + MWS("super-secret-token", PROJECT_URL + "?secret=do-not-log") + + warning.assert_called_once_with( + "mws_model_discovery_validation_failed", + extra={ + "provider": "MWS", + "operation": "model_discovery", + "validation_target": "api_url", + "error_type": "ValueError", + }, + ) + logged = repr(warning.call_args) + assert "super-secret-token" not in logged + assert "do-not-log" not in logged + assert "Bearer" not in logged + + +@pytest.mark.p1 +@pytest.mark.asyncio +async def test_mws_model_list_uses_exact_url_and_bearer_header(): + """Load and classify models with the required URL and bearer token.""" + response = MagicMock(status=200) + response.json = AsyncMock( + return_value={ + "data": [ + {"id": "qwen3-32b"}, + {"id": "qwen-vl"}, + {"id": "bge-m3"}, + {"id": "bge-reranker-v2-m3"}, + ] + } + ) + session = MagicMock() + session.get.return_value = _async_context(response) + session_context = _async_context(session) + + with ( + patch( + "rag.llm.model_meta.aiohttp.ClientSession", + return_value=session_context, + ) as client_session, + patch("rag.llm.model_meta.logging.info") as info, + ): + models = await MWS("token", PROJECT_URL + "/").get_model_list() + + assert client_session.call_args.kwargs["timeout"].total == 30 + session.get.assert_called_once_with( + PROJECT_URL + "/openai/v1/models", + headers={"Authorization": "Bearer token"}, + ) + assert [model["model_types"] for model in models] == [ + ["chat"], + ["embedding"], + ["rerank"], + ] + log_context = { + "provider": "MWS", + "operation": "model_discovery", + "url": PROJECT_URL + "/openai/v1/models", + } + assert info.call_args_list == [ + call("mws_model_discovery_request", extra=log_context), + call( + "mws_model_discovery_completed", + extra={**log_context, "result_count": 3}, + ), + ] + assert "token" not in repr(info.call_args_list) + assert "Bearer" not in repr(info.call_args_list) + + +@pytest.mark.p1 +@pytest.mark.asyncio +async def test_mws_model_discovery_logs_http_failure_and_zero_results(): + """Log a failed HTTP response and the resulting empty model list.""" + response = MagicMock(status=503) + session = MagicMock() + session.get.return_value = _async_context(response) + log_context = { + "provider": "MWS", + "operation": "model_discovery", + "url": PROJECT_URL + "/openai/v1/models", + } + + with ( + patch( + "rag.llm.model_meta.aiohttp.ClientSession", + return_value=_async_context(session), + ), + patch("rag.llm.model_meta.logging.info") as info, + patch("rag.llm.model_meta.logging.warning") as warning, + ): + models = await MWS("super-secret-token", PROJECT_URL).get_model_list() + + assert models == [] + warning.assert_called_once_with( + "mws_model_discovery_request_failed", + extra={ + **log_context, + "failure_stage": "http_response", + "http_status": 503, + }, + ) + assert info.call_args_list == [ + call("mws_model_discovery_request", extra=log_context), + call( + "mws_model_discovery_completed", + extra={**log_context, "result_count": 0}, + ), + ] + logged = repr([info.call_args_list, warning.call_args_list]) + assert "super-secret-token" not in logged + assert "Bearer" not in logged + + +@pytest.mark.p1 +@pytest.mark.asyncio +async def test_mws_model_discovery_logs_request_exception_without_credentials(): + """Log a discovery transport exception without exposing authentication.""" + session = MagicMock() + session.get.side_effect = RuntimeError("connection failed") + + with ( + patch( + "rag.llm.model_meta.aiohttp.ClientSession", + return_value=_async_context(session), + ), + patch("rag.llm.model_meta.logging.info"), + patch("rag.llm.model_meta.logging.warning") as warning, + ): + with pytest.raises(RuntimeError, match="connection failed"): + await MWS("super-secret-token", PROJECT_URL).get_model_list() + + warning.assert_called_once_with( + "mws_model_discovery_request_failed", + extra={ + "provider": "MWS", + "operation": "model_discovery", + "url": PROJECT_URL + "/openai/v1/models", + "failure_stage": "request", + "error_type": "RuntimeError", + }, + ) + logged = repr(warning.call_args) + assert "super-secret-token" not in logged + assert "Bearer" not in logged + + +@pytest.mark.p1 +@pytest.mark.asyncio +async def test_mws_chat_uses_exact_url_bearer_header_and_documented_fields(): + """Send only documented fields in a non-streaming MWS chat request.""" + chat = MWSChat("token", "qwen3-32b", PROJECT_URL + "/") + response = MagicMock(status=200) + response.json = AsyncMock( + return_value={ + "id": "chat-1", + "model": "qwen3-32b", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "Hello"}, + } + ], + "usage": { + "prompt_tokens": 4, + "completion_tokens": 1, + "total_tokens": 5, + }, + } + ) + session = MagicMock() + session.post.return_value = _async_context(response) + + with patch( + "rag.llm.chat_model.aiohttp.ClientSession", + return_value=_async_context(session), + ): + answer, tokens = await chat._async_chat( + [ + {"role": "system", "content": "Be concise."}, + { + "role": "user", + "content": "Hi", + "tool_call_id": "must-not-be-forwarded", + }, + ], + { + "temperature": 0.25, + "max_tokens": 128, + "top_p": 0.9, + "tools": [{"must": "not be forwarded"}], + }, + ) + + assert answer == "Hello" + assert tokens == 5 + assert chat.last_usage == { + "prompt_tokens": 4, + "completion_tokens": 1, + "total_tokens": 5, + } + session.post.assert_called_once_with( + PROJECT_URL + "/openai/v1/chat/completions", + headers={ + "Content-Type": "application/json", + "Authorization": "Bearer token", + }, + json={ + "model": "qwen3-32b", + "messages": [ + {"role": "system", "content": "Be concise."}, + {"role": "user", "content": "Hi"}, + ], + "temperature": 0.25, + "max_completion_tokens": 128, + }, + ) + + +@pytest.mark.p1 +@pytest.mark.asyncio +async def test_mws_chat_streaming_uses_documented_fields_and_usage(): + """Parse streaming MWS chat chunks and report final token usage.""" + chat = MWSChat("token", "qwen3-32b", PROJECT_URL) + response = MagicMock(status=200) + response.content = MagicMock() + response.content.__aiter__.return_value = iter( + [ + b'data: {"model":"qwen3-32b","choices":[{"index":0,"delta":{"content":"Hel"}}]}\n', + b'data: {"model":"qwen3-32b","choices":[{"index":0,"delta":{"content":"lo"},"finish_reason":"stop"}],"usage":{"prompt_tokens":4,"completion_tokens":1,"total_tokens":5}}\n', + b"data: [DONE]\n", + ] + ) + session = MagicMock() + session.post.return_value = _async_context(response) + + with patch( + "rag.llm.chat_model.aiohttp.ClientSession", + return_value=_async_context(session), + ): + chunks = [ + chunk + async for chunk in chat._async_chat_streamly( + [{"role": "user", "content": "Hi"}], + {"top_p": 0.9}, + ) + ] + + assert chunks == [("Hel", 0), ("lo", 5)] + assert chat.last_usage == { + "prompt_tokens": 4, + "completion_tokens": 1, + "total_tokens": 5, + } + session.post.assert_called_once_with( + PROJECT_URL + "/openai/v1/chat/completions", + headers={ + "Content-Type": "application/json", + "Authorization": "Bearer token", + }, + json={ + "model": "qwen3-32b", + "messages": [{"role": "user", "content": "Hi"}], + "stream": True, + "stream_options": {"include_usage": True}, + }, + ) + + +@pytest.mark.p1 +def test_mws_embedding_sends_only_documented_fields_and_orders_results(): + """Send a strict embedding request and restore response index order.""" + embed = MWSEmbed("token", "bge-m3", PROJECT_URL + "/") + response = _response( + { + "data": [ + {"index": 1, "embedding": [0.3, 0.4]}, + {"index": 0, "embedding": [0.1, 0.2]}, + ], + "usage": {"prompt_tokens": 7, "total_tokens": 7}, + } + ) + + with patch("rag.llm.embedding_model.requests.post", return_value=response) as post: + vectors, tokens = embed.encode(["first", "second"]) + + assert np.array_equal(vectors, np.array([[0.1, 0.2], [0.3, 0.4]])) + assert tokens == 7 + post.assert_called_once_with( + PROJECT_URL + "/openai/v1/embeddings", + headers={"Content-Type": "application/json", "Authorization": "Bearer token"}, + json={"model": "bge-m3", "input": ["first", "second"]}, + timeout=30, + ) + + +@pytest.mark.p1 +def test_mws_embedding_rejects_empty_token_without_request(): + """Reject an empty MWS token before attempting an embedding request.""" + with patch("rag.llm.embedding_model.requests.post") as post: + with pytest.raises(ValueError, match="Token is required"): + MWSEmbed(" ", "bge-m3", PROJECT_URL) + post.assert_not_called() + + +@pytest.mark.p1 +def test_mws_rerank_uses_cohere_path_and_exact_payload(): + """Use the MWS Cohere endpoint with the exact reranking payload.""" + reranker = MWSRerank("super-secret-token", "bge-reranker-v2-m3", PROJECT_URL) + response = _response( + { + "results": [ + {"index": 1, "relevance_score": 0.9}, + {"index": 0, "relevance_score": 0.2}, + ] + } + ) + + token_counts = {"query": 5, "first": 2, "second": 3} + with ( + patch("rag.llm.rerank_model.requests.post", return_value=response) as post, + patch( + "rag.llm.rerank_model.num_tokens_from_string", + side_effect=lambda text: token_counts[text], + ), + patch("rag.llm.rerank_model.logging.info") as info, + ): + scores, tokens = reranker.similarity("query", ["first", "second"]) + + assert np.array_equal(scores, np.array([0.2, 0.9])) + assert tokens == 10 + post.assert_called_once_with( + PROJECT_URL + "/cohere/v2/rerank", + headers={ + "Content-Type": "application/json", + "Authorization": "Bearer super-secret-token", + }, + json={ + "model": "bge-reranker-v2-m3", + "query": "query", + "documents": ["first", "second"], + "top_n": 2, + }, + timeout=30, + ) + info.assert_called_once_with( + "mws_rerank_request", + extra={ + "provider": "MWS", + "operation": "rerank", + "endpoint": PROJECT_URL + "/cohere/v2/rerank", + "model": "bge-reranker-v2-m3", + "document_count": 2, + }, + ) + logged = repr(info.call_args) + assert "super-secret-token" not in logged + assert "Bearer" not in logged + assert "query" not in logged + assert "first" not in logged + assert "second" not in logged + + +@pytest.mark.p1 +@pytest.mark.parametrize( + "result", + [ + {"index": 0}, + {"index": 0, "relevance_score": "invalid"}, + {"index": 0, "relevance_score": True}, + {"index": 0, "relevance_score": float("nan")}, + {"index": 0, "relevance_score": float("inf")}, + ], +) +def test_mws_rerank_rejects_invalid_relevance_scores(result): + """Reject missing, non-numeric, boolean, and non-finite scores.""" + reranker = MWSRerank("super-secret-token", "bge-reranker-v2-m3", PROJECT_URL) + response = _response({"results": [result]}) + log_context = { + "provider": "MWS", + "operation": "rerank", + "endpoint": PROJECT_URL + "/cohere/v2/rerank", + "model": "bge-reranker-v2-m3", + "document_count": 1, + } + + with ( + patch("rag.llm.rerank_model.requests.post", return_value=response), + patch("rag.llm.rerank_model.logging.info"), + patch("rag.llm.rerank_model.logging.warning") as warning, + ): + with pytest.raises(ValueError, match="relevance_score"): + reranker.similarity("query", ["first"]) + + warning.assert_called_once_with( + "mws_rerank_failed", + extra={ + **log_context, + "failure_stage": "response_validation", + "error_type": "ValueError", + }, + ) + logged = repr(warning.call_args) + assert "super-secret-token" not in logged + assert "query" not in logged + assert "first" not in logged + + +@pytest.mark.p1 +@pytest.mark.parametrize("failure_stage", ["http_request", "json_parsing"]) +def test_mws_rerank_logs_failure_stage_and_reraises(failure_stage): + """Log the failing rerank stage and re-raise the original exception.""" + reranker = MWSRerank("super-secret-token", "bge-reranker-v2-m3", PROJECT_URL) + response = _response({"results": []}) + error = RuntimeError(failure_stage) + if failure_stage == "http_request": + response.raise_for_status.side_effect = error + else: + response.json.side_effect = error + log_context = { + "provider": "MWS", + "operation": "rerank", + "endpoint": PROJECT_URL + "/cohere/v2/rerank", + "model": "bge-reranker-v2-m3", + "document_count": 1, + } + + with ( + patch("rag.llm.rerank_model.requests.post", return_value=response), + patch("rag.llm.rerank_model.logging.info"), + patch("rag.llm.rerank_model.logging.warning") as warning, + ): + with pytest.raises(RuntimeError) as caught: + reranker.similarity("query", ["first"]) + + assert caught.value is error + warning.assert_called_once_with( + "mws_rerank_failed", + extra={ + **log_context, + "failure_stage": failure_stage, + "error_type": "RuntimeError", + }, + ) + logged = repr(warning.call_args) + assert "super-secret-token" not in logged + assert "query" not in logged + assert "first" not in logged + + +@pytest.mark.p1 +def test_mws_empty_inputs_do_not_send_requests(): + """Return empty results without calling MWS for empty input lists.""" + embed = MWSEmbed("token", "bge-m3", PROJECT_URL) + reranker = MWSRerank("token", "bge-reranker-v2-m3", PROJECT_URL) + + with ( + patch("rag.llm.embedding_model.requests.post") as embed_post, + patch("rag.llm.rerank_model.requests.post") as rerank_post, + ): + vectors, embedding_tokens = embed.encode([]) + scores, rerank_tokens = reranker.similarity("query", []) + + assert vectors.size == 0 + assert scores.size == 0 + assert embedding_tokens == 0 + assert rerank_tokens == 0 + embed_post.assert_not_called() + rerank_post.assert_not_called() + + +@pytest.mark.p1 +def test_mws_rejects_incomplete_indexed_responses(): + """Reject incomplete embedding and reranking response collections.""" + embed = MWSEmbed("token", "bge-m3", PROJECT_URL) + reranker = MWSRerank("token", "bge-reranker-v2-m3", PROJECT_URL) + + with patch( + "rag.llm.embedding_model.requests.post", + return_value=_response({"data": [{"index": 0, "embedding": [0.1]}]}), + ): + with pytest.raises(Exception, match="1 embeddings for 2 inputs"): + embed.encode(["first", "second"]) + + with patch( + "rag.llm.rerank_model.requests.post", + return_value=_response({"results": [{"index": 0, "relevance_score": 0.1}]}), + ): + with pytest.raises(ValueError, match="1 rerank results for 2 documents"): + reranker.similarity("query", ["first", "second"]) + + +@pytest.mark.p1 +def test_mws_model_list_keeps_chat_embedding_and_rerank(): + """Expose only the three model types implemented by the MWS provider.""" + meta = MWS("token", PROJECT_URL) + models = meta._format_model_list( + { + "data": [ + {"id": "bge-m3"}, + {"id": "bge-reranker-v2-m3"}, + {"id": "qwen3-32b"}, + {"id": "qwen-vl"}, + ] + } + ) + + assert meta._get_model_list_url() == PROJECT_URL + "/openai/v1/models" + assert models == [ + { + "name": "bge-m3", + "model_types": ["embedding"], + "features": [], + "max_tokens": 8192, + }, + { + "name": "bge-reranker-v2-m3", + "model_types": ["rerank"], + "features": [], + "max_tokens": 8192, + }, + { + "name": "qwen3-32b", + "model_types": ["chat"], + "features": [], + "max_tokens": 8192, + }, + ] diff --git a/test/unit_test/rag/llm/test_rerank_normalization.py b/test/unit_test/rag/llm/test_rerank_normalization.py index 5af9abc467..26a1649a12 100644 --- a/test/unit_test/rag/llm/test_rerank_normalization.py +++ b/test/unit_test/rag/llm/test_rerank_normalization.py @@ -147,6 +147,47 @@ def test_nvidia_logits_are_normalized(): assert rank.min() >= 0.0 and rank.max() <= 1.0 +@pytest.mark.parametrize( + "configured_url", + [ + "https://ai.example.com/v1/retrieval/nvidia", + "https://ai.example.com/v1/retrieval/nvidia/", + ], +) +def test_nvidia_fallback_preserves_path_and_logs_without_api_key(configured_url): + """Preserve the base path and log fallback metadata without credentials.""" + with patch("rag.llm.rerank_model.logging.info") as info: + reranker = NvidiaRerank( + "super-secret-key", + "nvidia/custom-reranker", + configured_url, + ) + + expected_endpoint = "https://ai.example.com/v1/retrieval/nvidia/reranking" + assert reranker.base_url == expected_endpoint + info.assert_called_once_with( + "nvidia_rerank_fallback_endpoint_assigned", + extra={ + "provider": "NVIDIA", + "model": "nvidia/custom-reranker", + "endpoint": expected_endpoint, + }, + ) + assert "super-secret-key" not in repr(info.call_args) + + +def test_nvidia_specific_model_does_not_log_fallback_assignment(): + """Use the model-specific endpoint without emitting a fallback event.""" + with patch("rag.llm.rerank_model.logging.info") as info: + reranker = NvidiaRerank( + "key", + "nvidia/nv-rerankqa-mistral-4b-v3", + ) + + assert reranker.base_url == "https://ai.api.nvidia.com/v1/retrieval/nvidia/nv-rerankqa-mistral-4b-v3/reranking" + info.assert_not_called() + + def test_calibrated_relevance_scores_are_preserved(): # A provider already returning [0,1] relevance scores keeps them verbatim; # min-max would have stretched these to [1.0, 0.0, 0.5]. diff --git a/web/src/assets/svg/llm/mws.svg b/web/src/assets/svg/llm/mws.svg new file mode 100644 index 0000000000..1ab5b53ec2 --- /dev/null +++ b/web/src/assets/svg/llm/mws.svg @@ -0,0 +1,4 @@ + + + + diff --git a/web/src/components/svg-icon.tsx b/web/src/components/svg-icon.tsx index ab8e3f0baf..33dc7ffb23 100644 --- a/web/src/components/svg-icon.tsx +++ b/web/src/components/svg-icon.tsx @@ -126,6 +126,7 @@ const svgIcons = [ LLMFactory.FunASR, LLMFactory.AIMLAPI, LLMFactory.GreenPT, + LLMFactory.MWS, ]; export const LlmIcon = ({ diff --git a/web/src/constants/llm.ts b/web/src/constants/llm.ts index 5e8e86bad2..9218b43f1e 100644 --- a/web/src/constants/llm.ts +++ b/web/src/constants/llm.ts @@ -44,6 +44,7 @@ export enum LLMFactory { NVIDIA = 'NVIDIA', LMStudio = 'LM-Studio', OpenAiAPICompatible = 'OpenAI-API-Compatible', + MWS = 'MWS', Cohere = 'Cohere', LeptonAI = 'LeptonAI', TogetherAI = 'TogetherAI', @@ -193,6 +194,7 @@ export const IconMap = { [LLMFactory.FunASR]: 'funasr', [LLMFactory.AIMLAPI]: 'aimlapi', [LLMFactory.GreenPT]: 'greenpt', + [LLMFactory.MWS]: 'mws', }; export const ModelTypeToField: Record = { @@ -217,6 +219,8 @@ export const APIMapUrl = { [LLMFactory.OpenAI]: 'https://platform.openai.com/api-keys', [LLMFactory.AIMLAPI]: 'https://aimlapi.com/app/keys', [LLMFactory.GreenPT]: 'https://greenpt.ai', + [LLMFactory.MWS]: + 'https://mws.ru/docs/cloud-platform/gpt/general/inference-text.html', [LLMFactory.Anthropic]: 'https://console.anthropic.com/settings/keys', [LLMFactory.Gemini]: 'https://aistudio.google.com/app/apikey', [LLMFactory.DeepSeek]: 'https://platform.deepseek.com/api_keys', diff --git a/web/src/locales/en.ts b/web/src/locales/en.ts index 47f7299052..29eafe2f96 100644 --- a/web/src/locales/en.ts +++ b/web/src/locales/en.ts @@ -2154,6 +2154,13 @@ Example: Virtual Hosted Style`, modelTypeMessage: 'Please input your model type!', addLlmBaseUrl: 'Base URL', baseUrlNameMessage: 'Please input your Base URL', + mwsApiUrl: 'API URL', + mwsApiUrlMessage: 'Please enter the MWS project API URL', + mwsApiUrlPlaceholder: + 'https://gpt.mwsapis.ru/projects/', + mwsToken: 'Token', + mwsTokenMessage: 'Please enter the MWS Token', + mwsTokenPlaceholder: 'MWS service account API key', paddleocr: { apiUrl: 'PaddleOCR API URL', apiUrlPlaceholder: diff --git a/web/src/locales/ru.ts b/web/src/locales/ru.ts index 00ca6e804b..ce8bb25532 100644 --- a/web/src/locales/ru.ts +++ b/web/src/locales/ru.ts @@ -1301,6 +1301,13 @@ export default { modelTypeMessage: 'Пожалуйста, введите тип вашей модели!', addLlmBaseUrl: 'Базовый URL', baseUrlNameMessage: 'Пожалуйста, введите ваш базовый URL', + mwsApiUrl: 'URL API', + mwsApiUrlMessage: 'Введите URL API проекта MWS', + mwsApiUrlPlaceholder: + 'https://gpt.mwsapis.ru/projects/', + mwsToken: 'Токен', + mwsTokenMessage: 'Введите токен MWS', + mwsTokenPlaceholder: 'API-ключ сервисного аккаунта MWS', paddleocr: { apiUrl: 'URL API PaddleOCR', apiUrlPlaceholder: diff --git a/web/src/pages/user-setting/setting-model/provider-schema/constants.ts b/web/src/pages/user-setting/setting-model/provider-schema/constants.ts index 5cb60ca15b..2ea2f4e5b0 100644 --- a/web/src/pages/user-setting/setting-model/provider-schema/constants.ts +++ b/web/src/pages/user-setting/setting-model/provider-schema/constants.ts @@ -38,6 +38,7 @@ export const LIST_MODEL_PROVIDERS = new Set([ LLMFactory.OpenRouter, LLMFactory.VLLM, LLMFactory.OpenAiAPICompatible, + LLMFactory.MWS, LLMFactory.LMStudio, LLMFactory.VolcEngine, LLMFactory.Xinference, diff --git a/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.test.ts b/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.test.ts index 9bcc67f0bc..fd13728d84 100644 --- a/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.test.ts +++ b/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.test.ts @@ -36,3 +36,46 @@ describe('FunASR local provider configuration', () => { }); }); }); + +describe('MWS provider configuration', () => { + it('uses dynamic model discovery', () => { + expect(LIST_MODEL_PROVIDERS.has(LLMFactory.MWS)).toBe(true); + }); + + it('requires a project API URL and Token', () => { + const config = LocalLlmConfigs[LLMFactory.MWS]; + + expect(config).toMatchObject({ + llmFactory: LLMFactory.MWS, + title: 'MWS', + docLink: + 'https://mws.ru/docs/cloud-platform/gpt/general/inference-text.html', + }); + expect( + config.fields.find((field) => field.name === 'instance_name'), + ).toMatchObject({ label: 'instanceName', required: true }); + expect( + config.fields.find((field) => field.name === 'base_url'), + ).toMatchObject({ label: 'mwsApiUrl', required: true }); + expect( + config.fields.find((field) => field.name === 'api_key'), + ).toMatchObject({ label: 'mwsToken', required: true }); + expect( + config.fields.some((field) => field.name === 'provider_order'), + ).toBe(false); + + expect( + config.submitTransform?.({ + instance_name: 'mws-project', + base_url: 'https://gpt.mwsapis.ru/projects/demo', + api_key: 'token', + model_info: [], + }), + ).toMatchObject({ + instance_name: 'mws-project', + llm_factory: LLMFactory.MWS, + base_url: 'https://gpt.mwsapis.ru/projects/demo', + api_key: 'token', + }); + }); +}); diff --git a/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.ts b/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.ts index 38fad18a59..5c5e85540f 100644 --- a/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.ts +++ b/web/src/pages/user-setting/setting-model/provider-schema/field-config/local-llm-configs.ts @@ -104,6 +104,34 @@ export const LocalLlmConfigs: Record = { undefined, 'https://platform.openai.com/docs/models/gpt-4', ), + [LLMFactory.MWS]: buildLocalConfig( + LLMFactory.MWS, + 'MWS', + false, + [ + { + name: 'base_url', + label: 'mwsApiUrl', + type: 'inputSelect', + required: true, + placeholder: 'mwsApiUrlPlaceholder', + autoComplete: 'new-password', + shouldRender: 'hideWhenInstanceExists', + validation: { message: 'mwsApiUrlMessage' }, + }, + { + name: 'api_key', + label: 'mwsToken', + type: FormFieldType.Password, + required: true, + placeholder: 'mwsTokenPlaceholder', + autoComplete: 'new-password', + shouldRender: 'hideWhenInstanceExists', + validation: { message: 'mwsTokenMessage' }, + }, + ], + 'https://mws.ru/docs/cloud-platform/gpt/general/inference-text.html', + ), [LLMFactory.RAGcon]: buildLocalConfig( LLMFactory.RAGcon, 'RAGcon',