From 8bd5768ebcdd798ced7c5223489499a4c3974a52 Mon Sep 17 00:00:00 2001 From: nikminer Date: Tue, 11 Aug 2026 14:12:42 +0300 Subject: [PATCH] Integrate MWS model with API support and enhance chat functionality (#17959) ## What This pull request adds **MWS GPT Model Hub** as a built-in model provider in RAGFlow. The integration allows users to configure an MWS project endpoint and token, discover the models available to that project, and use supported MWS models for chat completion, embeddings, and reranking. Co-authored-by: ilarionov_n --- conf/llm_factories.json | 8 + conf/models/mws.json | 15 + docs/guides/models/supported_models.mdx | 1 + internal/entity/models/factory.go | 2 + internal/entity/models/mws.go | 316 +++++++++ internal/entity/models/mws_test.go | 326 ++++++++++ rag/llm/chat_model.py | 133 ++++ rag/llm/embedding_model.py | 36 ++ rag/llm/model_meta.py | 126 ++++ rag/llm/mws_utils.py | 43 ++ rag/llm/rerank_model.py | 106 ++- test/unit_test/rag/llm/test_mws.py | 607 ++++++++++++++++++ .../rag/llm/test_rerank_normalization.py | 41 ++ web/src/assets/svg/llm/mws.svg | 4 + web/src/components/svg-icon.tsx | 1 + web/src/constants/llm.ts | 4 + web/src/locales/en.ts | 7 + web/src/locales/ru.ts | 7 + .../provider-schema/constants.ts | 1 + .../field-config/local-llm-configs.test.ts | 43 ++ .../field-config/local-llm-configs.ts | 28 + 21 files changed, 1853 insertions(+), 2 deletions(-) create mode 100644 conf/models/mws.json create mode 100644 internal/entity/models/mws.go create mode 100644 internal/entity/models/mws_test.go create mode 100644 rag/llm/mws_utils.py create mode 100644 test/unit_test/rag/llm/test_mws.py create mode 100644 web/src/assets/svg/llm/mws.svg 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',