mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-12 20:03:40 +08:00
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:
@@ -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":
|
||||
|
||||
316
internal/entity/models/mws.go
Normal file
316
internal/entity/models/mws.go
Normal 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
|
||||
}
|
||||
326
internal/entity/models/mws_test.go
Normal file
326
internal/entity/models/mws_test.go
Normal 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())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user