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 <ilarionov_n@promis.ru>
This commit is contained in:
nikminer
2026-08-11 14:12:42 +03:00
committed by GitHub
parent 5e2c0eee28
commit 8bd5768ebc
21 changed files with 1853 additions and 2 deletions

View File

@@ -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":

View File

@@ -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/<project>")
}
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/<project>")
}
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
}

View File

@@ -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())
}
}