From c27ca82010cdf2b490298ba75abd6a1ecd785d6c Mon Sep 17 00:00:00 2001 From: Jerry <1736355688@qq.com> Date: Tue, 4 Aug 2026 17:57:42 +0800 Subject: [PATCH 1/2] feat(vertex): support OpenAI-compatible embeddings --- relay/channel/vertex/adaptor.go | 81 ++++++++++- relay/channel/vertex/dto.go | 31 +++++ relay/channel/vertex/embedding.go | 64 +++++++++ relay/channel/vertex/embedding_test.go | 180 +++++++++++++++++++++++++ 4 files changed, 354 insertions(+), 2 deletions(-) create mode 100644 relay/channel/vertex/embedding.go create mode 100644 relay/channel/vertex/embedding_test.go diff --git a/relay/channel/vertex/adaptor.go b/relay/channel/vertex/adaptor.go index c60d75d29f2..e5746d621e5 100644 --- a/relay/channel/vertex/adaptor.go +++ b/relay/channel/vertex/adaptor.go @@ -50,6 +50,8 @@ var claudeModelMap = map[string]string{ const anthropicVersion = "vertex-2023-10-16" +const geminiEmbedding001MaxDimensions = 3072 + type Adaptor struct { RequestMode int AccountCredentials Credentials @@ -170,6 +172,10 @@ func (a *Adaptor) getRequestUrl(info *relaycommon.RelayInfo, modelName, suffix s func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) { suffix := "" if a.RequestMode == RequestModeGemini { + if info.RelayMode == constant.RelayModeEmbeddings { + return a.getRequestUrl(info, info.UpstreamModelName, "predict") + } + if model_setting.GetGeminiSettings().ThinkingAdapterEnabled && !model_setting.ShouldPreserveThinkingSuffix(info.OriginModelName) { // 新增逻辑:处理 -thinking- 格式 @@ -323,8 +329,75 @@ func (a *Adaptor) ConvertRerankRequest(c *gin.Context, relayMode int, request dt } func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.EmbeddingRequest) (any, error) { - //TODO implement me - return nil, errors.New("not implemented") + input, inputErr := parseSingleVertexEmbeddingInput(request.Input) + if inputErr != nil { + return nil, invalidVertexEmbeddingRequest(inputErr) + } + if request.EncodingFormat != "" && request.EncodingFormat != "float" { + return nil, invalidVertexEmbeddingRequest( + fmt.Errorf("Vertex embedding does not support encoding_format %q", request.EncodingFormat), + ) + } + if request.Dimensions != nil && *request.Dimensions <= 0 { + return nil, invalidVertexEmbeddingRequest(errors.New("Vertex embedding dimensions must be greater than zero")) + } + if request.Dimensions != nil && info.UpstreamModelName == "gemini-embedding-001" && *request.Dimensions > geminiEmbedding001MaxDimensions { + return nil, invalidVertexEmbeddingRequest( + fmt.Errorf( + "Vertex model %s supports at most %d embedding dimensions", + info.UpstreamModelName, + geminiEmbedding001MaxDimensions, + ), + ) + } + + vertexRequest := &VertexEmbeddingRequest{ + Instances: []VertexEmbeddingInstance{{Content: input}}, + } + if request.Dimensions != nil { + vertexRequest.Parameters = &VertexEmbeddingParameters{ + OutputDimensionality: request.Dimensions, + } + } + return vertexRequest, nil +} + +func parseSingleVertexEmbeddingInput(input any) (string, error) { + var value any + switch typedInput := input.(type) { + case string: + value = typedInput + case []any: + if len(typedInput) != 1 { + return "", fmt.Errorf("Vertex embedding requires exactly one input, got %d", len(typedInput)) + } + value = typedInput[0] + case []string: + if len(typedInput) != 1 { + return "", fmt.Errorf("Vertex embedding requires exactly one input, got %d", len(typedInput)) + } + value = typedInput[0] + default: + return "", fmt.Errorf("Vertex embedding input must be a string or a single-element string array, got %T", input) + } + + text, ok := value.(string) + if !ok { + return "", fmt.Errorf("Vertex embedding input array must contain a string, got %T", value) + } + if strings.TrimSpace(text) == "" { + return "", errors.New("Vertex embedding input must not be empty") + } + return text, nil +} + +func invalidVertexEmbeddingRequest(err error) *types.NewAPIError { + return types.NewErrorWithStatusCode( + err, + types.ErrorCodeInvalidRequest, + http.StatusBadRequest, + types.ErrOptionWithSkipRetry(), + ) } func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) { @@ -337,6 +410,10 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request } func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) { + if info.RelayMode == constant.RelayModeEmbeddings { + return VertexEmbeddingHandler(c, info, resp) + } + claudeAdaptor := claude.Adaptor{} if info.IsStream { switch a.RequestMode { diff --git a/relay/channel/vertex/dto.go b/relay/channel/vertex/dto.go index 7beab9ed0c6..978315140ea 100644 --- a/relay/channel/vertex/dto.go +++ b/relay/channel/vertex/dto.go @@ -23,6 +23,37 @@ type VertexAIClaudeRequest struct { //Metadata json.RawMessage `json:"metadata,omitempty"` } +type VertexEmbeddingInstance struct { + Content string `json:"content"` +} + +type VertexEmbeddingParameters struct { + OutputDimensionality *int `json:"outputDimensionality,omitempty"` +} + +type VertexEmbeddingRequest struct { + Instances []VertexEmbeddingInstance `json:"instances"` + Parameters *VertexEmbeddingParameters `json:"parameters,omitempty"` +} + +type VertexEmbeddingStatistics struct { + TokenCount int `json:"token_count"` + Truncated bool `json:"truncated"` +} + +type VertexEmbedding struct { + Values []float64 `json:"values"` + Statistics VertexEmbeddingStatistics `json:"statistics"` +} + +type VertexEmbeddingPrediction struct { + Embeddings VertexEmbedding `json:"embeddings"` +} + +type VertexEmbeddingResponse struct { + Predictions []VertexEmbeddingPrediction `json:"predictions"` +} + func copyRequest(req *dto.ClaudeRequest, version string) *VertexAIClaudeRequest { return &VertexAIClaudeRequest{ AnthropicVersion: version, diff --git a/relay/channel/vertex/embedding.go b/relay/channel/vertex/embedding.go new file mode 100644 index 00000000000..344a880538d --- /dev/null +++ b/relay/channel/vertex/embedding.go @@ -0,0 +1,64 @@ +package vertex + +import ( + "errors" + "io" + "net/http" + + "github.com/QuantumNous/new-api/common" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/relaykit/dto" + "github.com/QuantumNous/new-api/relaykit/types" + "github.com/QuantumNous/new-api/service" + + "github.com/gin-gonic/gin" +) + +func VertexEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { + defer service.CloseResponseBodyGracefully(resp) + + responseBody, err := io.ReadAll(resp.Body) + if err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + + var vertexResponse VertexEmbeddingResponse + if err := common.Unmarshal(responseBody, &vertexResponse); err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + if len(vertexResponse.Predictions) != 1 { + return nil, types.NewOpenAIError(errors.New("Vertex embedding response must contain exactly one prediction"), types.ErrorCodeBadResponseBody, http.StatusBadGateway) + } + + prediction := vertexResponse.Predictions[0] + if len(prediction.Embeddings.Values) == 0 { + return nil, types.NewOpenAIError(errors.New("Vertex embedding response contains an empty vector"), types.ErrorCodeBadResponseBody, http.StatusBadGateway) + } + + promptTokens := prediction.Embeddings.Statistics.TokenCount + if promptTokens == 0 { + promptTokens = info.GetEstimatePromptTokens() + } + usage := dto.Usage{ + PromptTokens: promptTokens, + TotalTokens: promptTokens, + InputTokens: promptTokens, + } + openAIResponse := dto.OpenAIEmbeddingResponse{ + Object: "list", + Data: []dto.OpenAIEmbeddingResponseItem{{ + Object: "embedding", + Index: 0, + Embedding: prediction.Embeddings.Values, + }}, + Model: info.UpstreamModelName, + Usage: usage, + } + + jsonResponse, err := common.Marshal(openAIResponse) + if err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError) + } + service.IOCopyBytesGracefully(c, resp, jsonResponse) + return &usage, nil +} diff --git a/relay/channel/vertex/embedding_test.go b/relay/channel/vertex/embedding_test.go new file mode 100644 index 00000000000..c6ffd3cc00e --- /dev/null +++ b/relay/channel/vertex/embedding_test.go @@ -0,0 +1,180 @@ +package vertex + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/QuantumNous/new-api/relaykit/dto" + "github.com/QuantumNous/new-api/relaykit/types" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestConvertEmbeddingRequestBuildsVertexPredictPayload(t *testing.T) { + t.Parallel() + + dimensions := 1024 + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: "gemini-embedding-001", + }, + } + request := dto.EmbeddingRequest{ + Input: []any{"hello Vertex"}, + EncodingFormat: "float", + Dimensions: &dimensions, + } + + converted, err := (&Adaptor{}).ConvertEmbeddingRequest(nil, info, request) + require.NoError(t, err) + require.Equal(t, &VertexEmbeddingRequest{ + Instances: []VertexEmbeddingInstance{{Content: "hello Vertex"}}, + Parameters: &VertexEmbeddingParameters{ + OutputDimensionality: &dimensions, + }, + }, converted) +} + +func TestConvertEmbeddingRequestRejectsUnsupportedInputs(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + request dto.EmbeddingRequest + }{ + { + name: "multiple inputs", + request: dto.EmbeddingRequest{Input: []any{"first", "second"}}, + }, + { + name: "mixed input array", + request: dto.EmbeddingRequest{Input: []any{123, "second"}}, + }, + { + name: "token id array", + request: dto.EmbeddingRequest{Input: []any{123}}, + }, + { + name: "empty input", + request: dto.EmbeddingRequest{Input: " "}, + }, + { + name: "base64 encoding", + request: dto.EmbeddingRequest{Input: "hello", EncodingFormat: "base64"}, + }, + { + name: "invalid dimensions", + request: dto.EmbeddingRequest{Input: "hello", Dimensions: intPointer(0)}, + }, + { + name: "dimensions above model limit", + request: dto.EmbeddingRequest{Input: "hello", Dimensions: intPointer(3073)}, + }, + } + + for _, test := range tests { + test := test + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: "gemini-embedding-001", + }, + } + converted, err := (&Adaptor{}).ConvertEmbeddingRequest(nil, info, test.request) + require.Nil(t, converted) + require.Error(t, err) + + var apiErr *types.NewAPIError + require.ErrorAs(t, err, &apiErr) + require.Equal(t, http.StatusBadRequest, apiErr.StatusCode) + require.Equal(t, types.ErrorCodeInvalidRequest, apiErr.GetErrorCode()) + }) + } +} + +func TestGetRequestURLUsesVertexPredictForEmbeddings(t *testing.T) { + t.Parallel() + + info := &relaycommon.RelayInfo{ + RelayMode: relayconstant.RelayModeEmbeddings, + OriginModelName: "gemini-embedding-001", + ChannelMeta: &relaycommon.ChannelMeta{ + ApiVersion: "us-central1", + ApiKey: `{"project_id":"airjelly-project"}`, + UpstreamModelName: "gemini-embedding-001", + }, + } + adaptor := &Adaptor{} + adaptor.Init(info) + + requestURL, err := adaptor.GetRequestURL(info) + require.NoError(t, err) + require.Equal( + t, + "https://us-central1-aiplatform.googleapis.com/v1/projects/airjelly-project/locations/us-central1/publishers/google/models/gemini-embedding-001:predict", + requestURL, + ) +} + +func TestVertexEmbeddingHandlerConvertsPredictResponse(t *testing.T) { + t.Parallel() + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + info := &relaycommon.RelayInfo{ + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: "gemini-embedding-001", + }, + } + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader( + `{"predictions":[{"embeddings":{"statistics":{"truncated":false,"token_count":6},"values":[0.25,-0.5]}}]}`, + )), + } + + usage, apiErr := VertexEmbeddingHandler(c, info, resp) + require.Nil(t, apiErr) + require.Equal(t, 6, usage.PromptTokens) + require.Equal(t, 6, usage.TotalTokens) + + var converted dto.OpenAIEmbeddingResponse + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &converted)) + require.Equal(t, "list", converted.Object) + require.Equal(t, "gemini-embedding-001", converted.Model) + require.Equal(t, []float64{0.25, -0.5}, converted.Data[0].Embedding) + require.Equal(t, 0, converted.Data[0].Index) + require.Equal(t, 6, converted.PromptTokens) +} + +func TestVertexEmbeddingHandlerRejectsEmptyPrediction(t *testing.T) { + t.Parallel() + + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + info := &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{}} + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(`{"predictions":[]}`)), + } + + usage, apiErr := VertexEmbeddingHandler(c, info, resp) + require.Nil(t, usage) + require.NotNil(t, apiErr) + require.Equal(t, http.StatusBadGateway, apiErr.StatusCode) +} + +func intPointer(value int) *int { + return &value +} From 553ec4f6c75dec706396adefed0a137c09bcbb4d Mon Sep 17 00:00:00 2001 From: Jerry <1736355688@qq.com> Date: Tue, 11 Aug 2026 19:21:54 +0800 Subject: [PATCH 2/2] fix(vertex): reject negative embedding usage --- relay/channel/vertex/embedding.go | 2 +- relay/channel/vertex/embedding_test.go | 27 ++++++++++++++++++++++++++ 2 files changed, 28 insertions(+), 1 deletion(-) diff --git a/relay/channel/vertex/embedding.go b/relay/channel/vertex/embedding.go index 344a880538d..739931e7645 100644 --- a/relay/channel/vertex/embedding.go +++ b/relay/channel/vertex/embedding.go @@ -36,7 +36,7 @@ func VertexEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *h } promptTokens := prediction.Embeddings.Statistics.TokenCount - if promptTokens == 0 { + if promptTokens <= 0 { promptTokens = info.GetEstimatePromptTokens() } usage := dto.Usage{ diff --git a/relay/channel/vertex/embedding_test.go b/relay/channel/vertex/embedding_test.go index c6ffd3cc00e..0c1a31ad54f 100644 --- a/relay/channel/vertex/embedding_test.go +++ b/relay/channel/vertex/embedding_test.go @@ -1,6 +1,7 @@ package vertex import ( + "fmt" "io" "net/http" "net/http/httptest" @@ -157,6 +158,32 @@ func TestVertexEmbeddingHandlerConvertsPredictResponse(t *testing.T) { require.Equal(t, 6, converted.PromptTokens) } +func TestVertexEmbeddingHandlerFallsBackForNonPositiveTokenCount(t *testing.T) { + t.Parallel() + + for _, tokenCount := range []int{0, -1} { + t.Run(fmt.Sprintf("token count %d", tokenCount), func(t *testing.T) { + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + info := &relaycommon.RelayInfo{ChannelMeta: &relaycommon.ChannelMeta{}} + info.SetEstimatePromptTokens(7) + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(fmt.Sprintf( + `{"predictions":[{"embeddings":{"statistics":{"token_count":%d},"values":[0.25]}}]}`, + tokenCount, + ))), + } + + usage, apiErr := VertexEmbeddingHandler(c, info, resp) + require.Nil(t, apiErr) + require.Equal(t, 7, usage.PromptTokens) + require.Equal(t, 7, usage.TotalTokens) + }) + } +} + func TestVertexEmbeddingHandlerRejectsEmptyPrediction(t *testing.T) { t.Parallel()