* test(relayconvert): add golden snapshot matrix and relaykit boundary guard Phase 0 of the relaykit extraction plan: pin byte-level output of every registered (from,to) request/response/stream conversion route, and forbid kit-bound packages from growing host-only imports. * wip(relayconvert): drop gin.Context from converter signatures; add convmeta draft Phase 1 in progress: relayconvert now takes context.Context; host media resolver adapts gin.Context back at the service boundary. * refactor(relayconvert): decouple converters from RelayInfo, gin, and settings Phase 1 of the relaykit extraction plan: - converters now depend on convmeta.Meta (implemented by RelayInfo) instead of *relaycommon.RelayInfo; ClaudeConvertInfo and the format guesser move to convmeta with aliases left behind - host settings reach converters via a convmeta.Options snapshot built in RelayInfo.ConvOptions; no more model_setting/reasoning global reads inside the conversion layer - effort-suffix helpers move to service/relayconvert/reasoning (old package forwards); chat-to-responses upgrade policy moves to service (host routing logic, not conversion) - golden conversion matrix unchanged * test(relayconvert): tighten boundary — kit packages now free of gin/setting imports * refactor(dto): drop gin and logger dependencies Phase 2 (part 1): dto.Request.IsStream now takes *http.Request instead of *gin.Context (Gemini's impl reads query/path off the std request); dto's three logger calls become common.SysError. Boundary test allowlist is now empty — kit-bound packages import no gin/setting/logger/model. * refactor(kit): extract dependency-free kitutil; dto/types/relayconvert stop importing common Phase 2 of the relaykit extraction plan: - new service/relayconvert/kitutil holds the pure helpers the kit needs (JSON wrappers, pointer/string/uuid/timestamp utils, MaskSensitiveInfo, pluggable LogInfo/LogError hooks, Debug flag) - dto, types, and all relayconvert packages now use kitutil; their only remaining internal deps are dto/types/constant - common keeps every original symbol (MaskSensitiveInfo delegates to kitutil) so host code is untouched; main.go routes kit logging into common.SysLog/SysError and mirrors DebugEnabled - golden conversion matrix unchanged * refactor(kit): move EndpointType/FinishReason to types; OpenRouter dialect via Options Kit packages (dto/types/relayconvert/reasonmap) no longer import constant: - EndpointType and finish-reason values live in types; constant re-exports - the OpenRouter special-case in claude->openai request conversion reads Options.OpenRouterDialect, set by the host from the channel type; InitChannelMeta invalidates the cached snapshot on channel switch * refactor: extract relaykit submodule (dto/types/relayconvert/reasonmap) Phase 3 of the relaykit extraction plan: - new go module github.com/QuantumNous/new-api/relaykit containing dto (minus task family), types, relayconvert (with convmeta/kitutil/reasoning), and reasonmap; host consumes it via require + replace, go.work for dev - task-family dto (task/suno/midjourney/video) stays in the host dto package; dual-consumer host files alias it as taskdto - relaykit builds and tests standalone (GOWORK=off): no host imports, no gin, no DB, no settings - golden conversion matrix unchanged * build(docker): copy relaykit/go.mod before go mod download The local-replace submodule's go.mod must exist inside the build context for the main module graph to resolve. * fix: address relaykit extraction regressions * fix: address relaykit review regressions * docs: document Meta nil receiver contract * fix(relaykit): fail OpenAI→Claude conversion without max_tokens; reject negative default_max_tokens The Claude Messages API requires max_tokens (omitting it is a 400 "Field required"), but with a nil Options.Claude.DefaultMaxTokens hook the converters silently emitted a request the upstream is guaranteed to reject. Both OpenAI Chat and Responses → Claude conversions now return sharedclaude.ErrMissingMaxTokens when no path (client value, default hook, thinking-adapter floor) supplied one. Unreachable in the host, which always configures the hook. Host side, claude.default_max_tokens now rejects negative values at the option API before persisting — they would wrap into huge unsigned values during conversion. Zero stays allowed: the current API treats max_tokens: 0 as cache pre-warming. * fix: make Gemini safety settings read path race-free
563 lines
19 KiB
Go
563 lines
19 KiB
Go
package gemini
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/constant"
|
|
"github.com/QuantumNous/new-api/logger"
|
|
"github.com/QuantumNous/new-api/relay/channel/openai"
|
|
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
|
"github.com/QuantumNous/new-api/relay/helper"
|
|
"github.com/QuantumNous/new-api/relaykit/dto"
|
|
"github.com/QuantumNous/new-api/relaykit/relayconvert"
|
|
"github.com/QuantumNous/new-api/relaykit/types"
|
|
"github.com/QuantumNous/new-api/service"
|
|
"github.com/gin-gonic/gin"
|
|
)
|
|
|
|
func buildUsageFromGeminiMetadata(metadata *dto.GeminiUsageMetadata, fallbackPromptTokens int) dto.Usage {
|
|
usage := relayconvert.UsageFromGeminiMetadata(metadata, fallbackPromptTokens)
|
|
if usage == nil {
|
|
return dto.Usage{}
|
|
}
|
|
return *usage
|
|
}
|
|
|
|
func attachEstimatedGeminiBillingUsage(usage *dto.Usage) *dto.Usage {
|
|
if usage != nil && usage.BillingUsage == nil {
|
|
usage.BillingUsage = dto.NewEstimatedGeminiChatBillingUsage(usage)
|
|
}
|
|
return usage
|
|
}
|
|
|
|
// patchGeminiZeroCompletionUsage estimates completion tokens locally when upstream
|
|
// usageMetadata was billable but reported zero completion tokens even though output
|
|
// content was actually received. Typical case: the client aborts a stream before the
|
|
// final chunk that carries candidatesTokenCount, leaving prompt-only metadata; without
|
|
// this patch the output side would settle at zero quota.
|
|
func patchGeminiZeroCompletionUsage(c *gin.Context, info *relaycommon.RelayInfo, usage *dto.Usage, responseText string, imageCount int) {
|
|
if usage == nil || usage.CompletionTokens > 0 {
|
|
return
|
|
}
|
|
if responseText == "" && imageCount == 0 {
|
|
return
|
|
}
|
|
estimated := service.ResponseText2Usage(c, responseText, info.UpstreamModelName, usage.PromptTokens)
|
|
usage.CompletionTokens = estimated.CompletionTokens
|
|
if imageCount != 0 && usage.CompletionTokens == 0 {
|
|
usage.CompletionTokens = imageCount * 1400
|
|
}
|
|
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
|
|
// Overwrite the metadata-derived billing usage: effectiveBillingUsage prefers
|
|
// BillingUsage during settlement, so keeping the prompt-only metadata there
|
|
// would still bill zero completion tokens.
|
|
usage.BillingUsage = dto.NewEstimatedGeminiChatBillingUsage(usage)
|
|
}
|
|
|
|
func geminiResponseUsageText(response *dto.GeminiChatResponse) string {
|
|
if response == nil {
|
|
return ""
|
|
}
|
|
var text strings.Builder
|
|
for _, candidate := range response.Candidates {
|
|
for _, part := range candidate.Content.Parts {
|
|
if part.Text != "" {
|
|
text.WriteString(part.Text)
|
|
}
|
|
}
|
|
}
|
|
return text.String()
|
|
}
|
|
|
|
func markGeminiGoogleSearchCall(c *gin.Context, response *dto.GeminiChatResponse) {
|
|
if c == nil || response == nil {
|
|
return
|
|
}
|
|
for _, candidate := range response.Candidates {
|
|
if candidate.GroundingMetadata != nil && len(candidate.GroundingMetadata.WebSearchQueries) > 0 {
|
|
c.Set("gemini_google_search_call", true)
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
func buildUsageFromGeminiResponse(c *gin.Context, info *relaycommon.RelayInfo, response *dto.GeminiChatResponse) dto.Usage {
|
|
metadata := response.GetUsageMetadata()
|
|
if dto.HasGeminiUsageMetadataTokens(metadata) {
|
|
usage := buildUsageFromGeminiMetadata(metadata, info.GetEstimatePromptTokens())
|
|
patchGeminiZeroCompletionUsage(c, info, &usage, geminiResponseUsageText(response), geminiResponseInlineImageCount(response))
|
|
return usage
|
|
}
|
|
usage := service.ResponseText2Usage(c, geminiResponseUsageText(response), info.UpstreamModelName, info.GetEstimatePromptTokens())
|
|
attachEstimatedGeminiBillingUsage(usage)
|
|
return *usage
|
|
}
|
|
|
|
func geminiResponseInlineImageCount(response *dto.GeminiChatResponse) int {
|
|
if response == nil {
|
|
return 0
|
|
}
|
|
count := 0
|
|
for _, candidate := range response.Candidates {
|
|
for _, part := range candidate.Content.Parts {
|
|
if part.InlineData != nil && part.InlineData.MimeType != "" {
|
|
count++
|
|
}
|
|
}
|
|
}
|
|
return count
|
|
}
|
|
|
|
func responseGeminiChat2OpenAI(c *gin.Context, response *dto.GeminiChatResponse) *dto.OpenAITextResponse {
|
|
return relayconvert.ResponseGeminiChat2OpenAI(helper.GetResponseID(c), common.GetTimestamp(), response)
|
|
}
|
|
|
|
func streamResponseGeminiChat2OpenAI(geminiResponse *dto.GeminiChatResponse) (*dto.ChatCompletionsStreamResponse, bool) {
|
|
return relayconvert.StreamResponseGeminiChat2OpenAI(geminiResponse)
|
|
}
|
|
|
|
func handleStream(c *gin.Context, info *relaycommon.RelayInfo, resp *dto.ChatCompletionsStreamResponse) error {
|
|
streamData, err := common.Marshal(resp)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to marshal stream response: %w", err)
|
|
}
|
|
err = openai.HandleStreamFormat(c, info, string(streamData), info.ChannelSetting.ForceFormat, info.ChannelSetting.ThinkingToContent)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to handle stream format: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func handleFinalStream(c *gin.Context, info *relaycommon.RelayInfo, resp *dto.ChatCompletionsStreamResponse) error {
|
|
streamData, err := common.Marshal(resp)
|
|
if err != nil {
|
|
return fmt.Errorf("failed to marshal stream response: %w", err)
|
|
}
|
|
openai.HandleFinalResponse(c, info, string(streamData), resp.Id, resp.Created, resp.Model, resp.GetSystemFingerprint(), resp.Usage, false)
|
|
return nil
|
|
}
|
|
|
|
func geminiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response, callback func(data string, geminiResponse *dto.GeminiChatResponse) bool) (*dto.Usage, *types.NewAPIError) {
|
|
var usage = &dto.Usage{}
|
|
var imageCount int
|
|
var hasBillableUsageMetadata bool
|
|
responseText := strings.Builder{}
|
|
|
|
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {
|
|
var geminiResponse dto.GeminiChatResponse
|
|
if err := common.UnmarshalJsonStr(data, &geminiResponse); err != nil {
|
|
sr.Stop(fmt.Errorf("unmarshal: %w", err))
|
|
return
|
|
}
|
|
|
|
if len(geminiResponse.Candidates) == 0 && geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil {
|
|
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason))
|
|
}
|
|
|
|
markGeminiGoogleSearchCall(c, &geminiResponse)
|
|
|
|
// 统计图片数量
|
|
for _, candidate := range geminiResponse.Candidates {
|
|
for _, part := range candidate.Content.Parts {
|
|
if part.InlineData != nil && part.InlineData.MimeType != "" {
|
|
imageCount++
|
|
}
|
|
if part.Text != "" {
|
|
responseText.WriteString(part.Text)
|
|
}
|
|
}
|
|
}
|
|
|
|
// 更新使用量统计
|
|
if metadata := geminiResponse.GetUsageMetadata(); dto.HasGeminiUsageMetadataTokens(metadata) {
|
|
mappedUsage := buildUsageFromGeminiMetadata(metadata, info.GetEstimatePromptTokens())
|
|
*usage = mappedUsage
|
|
hasBillableUsageMetadata = true
|
|
}
|
|
|
|
if !callback(data, &geminiResponse) {
|
|
sr.Stop(fmt.Errorf("gemini callback stopped"))
|
|
}
|
|
})
|
|
|
|
if !hasBillableUsageMetadata {
|
|
if info.ReceivedResponseCount > 0 {
|
|
usage = service.ResponseText2Usage(c, responseText.String(), info.UpstreamModelName, info.GetEstimatePromptTokens())
|
|
} else {
|
|
usage = &dto.Usage{}
|
|
}
|
|
if imageCount != 0 && usage.CompletionTokens == 0 {
|
|
usage.CompletionTokens = imageCount * 1400
|
|
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
|
|
common.SetContextKey(c, constant.ContextKeyLocalCountTokens, true)
|
|
}
|
|
attachEstimatedGeminiBillingUsage(usage)
|
|
} else {
|
|
patchGeminiZeroCompletionUsage(c, info, usage, responseText.String(), imageCount)
|
|
}
|
|
|
|
return usage, nil
|
|
}
|
|
|
|
func GeminiChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
|
|
id := helper.GetResponseID(c)
|
|
createAt := common.GetTimestamp()
|
|
finishReason := constant.FinishReasonStop
|
|
toolCallIndexByChoice := make(map[int]map[string]int)
|
|
nextToolCallIndexByChoice := make(map[int]int)
|
|
|
|
usage, err := geminiStreamHandler(c, info, resp, func(data string, geminiResponse *dto.GeminiChatResponse) bool {
|
|
response, isStop := streamResponseGeminiChat2OpenAI(geminiResponse)
|
|
|
|
response.Id = id
|
|
response.Created = createAt
|
|
response.Model = info.UpstreamModelName
|
|
if response.IsToolCall() {
|
|
finishReason = constant.FinishReasonToolCalls
|
|
if info.RelayFormat == types.RelayFormatClaude {
|
|
for choiceIdx := range response.Choices {
|
|
response.Choices[choiceIdx].FinishReason = nil
|
|
}
|
|
}
|
|
}
|
|
for choiceIdx := range response.Choices {
|
|
choiceKey := response.Choices[choiceIdx].Index
|
|
for toolIdx := range response.Choices[choiceIdx].Delta.ToolCalls {
|
|
tool := &response.Choices[choiceIdx].Delta.ToolCalls[toolIdx]
|
|
if tool.ID == "" {
|
|
continue
|
|
}
|
|
m := toolCallIndexByChoice[choiceKey]
|
|
if m == nil {
|
|
m = make(map[string]int)
|
|
toolCallIndexByChoice[choiceKey] = m
|
|
}
|
|
if idx, ok := m[tool.ID]; ok {
|
|
tool.SetIndex(idx)
|
|
continue
|
|
}
|
|
idx := nextToolCallIndexByChoice[choiceKey]
|
|
nextToolCallIndexByChoice[choiceKey] = idx + 1
|
|
m[tool.ID] = idx
|
|
tool.SetIndex(idx)
|
|
}
|
|
}
|
|
|
|
logger.LogDebug(c, "info.SendResponseCount = %d", info.SendResponseCount)
|
|
if info.SendResponseCount == 0 {
|
|
// send first response
|
|
emptyResponse := helper.GenerateStartEmptyResponse(id, createAt, info.UpstreamModelName, nil)
|
|
if response.IsToolCall() {
|
|
if len(emptyResponse.Choices) > 0 && len(response.Choices) > 0 {
|
|
toolCalls := response.Choices[0].Delta.ToolCalls
|
|
copiedToolCalls := make([]dto.ToolCallResponse, len(toolCalls))
|
|
for idx := range toolCalls {
|
|
copiedToolCalls[idx] = toolCalls[idx]
|
|
copiedToolCalls[idx].Function.Arguments = ""
|
|
}
|
|
emptyResponse.Choices[0].Delta.ToolCalls = copiedToolCalls
|
|
}
|
|
finishReason = constant.FinishReasonToolCalls
|
|
err := handleStream(c, info, emptyResponse)
|
|
if err != nil {
|
|
logger.LogError(c, err.Error())
|
|
}
|
|
|
|
response.ClearToolCalls()
|
|
if response.IsFinished() {
|
|
response.Choices[0].FinishReason = nil
|
|
}
|
|
} else {
|
|
err := handleStream(c, info, emptyResponse)
|
|
if err != nil {
|
|
logger.LogError(c, err.Error())
|
|
}
|
|
}
|
|
}
|
|
|
|
err := handleStream(c, info, response)
|
|
if err != nil {
|
|
logger.LogError(c, err.Error())
|
|
}
|
|
if isStop {
|
|
if info.RelayFormat != types.RelayFormatClaude {
|
|
_ = handleStream(c, info, helper.GenerateStopResponse(id, createAt, info.UpstreamModelName, finishReason))
|
|
}
|
|
}
|
|
return true
|
|
})
|
|
|
|
if err != nil {
|
|
return usage, err
|
|
}
|
|
|
|
response := helper.GenerateFinalUsageResponse(id, createAt, info.UpstreamModelName, *usage)
|
|
if info.RelayFormat == types.RelayFormatClaude && info.ClaudeConvertInfo != nil && !info.ClaudeConvertInfo.Done {
|
|
response = helper.GenerateStopResponse(id, createAt, info.UpstreamModelName, finishReason)
|
|
response.Usage = usage
|
|
}
|
|
handleErr := handleFinalStream(c, info, response)
|
|
if handleErr != nil {
|
|
common.SysLog("send final response failed: " + handleErr.Error())
|
|
}
|
|
return usage, nil
|
|
}
|
|
|
|
func GeminiChatHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
|
|
responseBody, err := io.ReadAll(resp.Body)
|
|
if err != nil {
|
|
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
service.CloseResponseBodyGracefully(resp)
|
|
logger.LogDebug(c, "Gemini response body: %s", responseBody)
|
|
var geminiResponse dto.GeminiChatResponse
|
|
err = common.Unmarshal(responseBody, &geminiResponse)
|
|
if err != nil {
|
|
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
markGeminiGoogleSearchCall(c, &geminiResponse)
|
|
if len(geminiResponse.Candidates) == 0 {
|
|
usage := buildUsageFromGeminiResponse(c, info, &geminiResponse)
|
|
|
|
var newAPIError *types.NewAPIError
|
|
if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil {
|
|
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason))
|
|
newAPIError = types.NewOpenAIError(
|
|
errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason),
|
|
types.ErrorCodePromptBlocked,
|
|
http.StatusBadRequest,
|
|
)
|
|
} else {
|
|
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "gemini_empty_candidates")
|
|
newAPIError = types.NewOpenAIError(
|
|
errors.New("empty response from Gemini API"),
|
|
types.ErrorCodeEmptyResponse,
|
|
http.StatusInternalServerError,
|
|
)
|
|
}
|
|
|
|
service.ResetStatusCode(newAPIError, c.GetString("status_code_mapping"))
|
|
|
|
switch info.RelayFormat {
|
|
case types.RelayFormatClaude:
|
|
c.JSON(newAPIError.StatusCode, gin.H{
|
|
"type": "error",
|
|
"error": newAPIError.ToClaudeError(),
|
|
})
|
|
default:
|
|
c.JSON(newAPIError.StatusCode, gin.H{
|
|
"error": newAPIError.ToOpenAIError(),
|
|
})
|
|
}
|
|
return &usage, nil
|
|
}
|
|
fullTextResponse := responseGeminiChat2OpenAI(c, &geminiResponse)
|
|
fullTextResponse.Model = info.UpstreamModelName
|
|
usage := buildUsageFromGeminiResponse(c, info, &geminiResponse)
|
|
|
|
fullTextResponse.Usage = usage
|
|
|
|
switch info.RelayFormat {
|
|
case types.RelayFormatOpenAI:
|
|
responseBody, err = common.Marshal(fullTextResponse)
|
|
if err != nil {
|
|
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
|
|
}
|
|
case types.RelayFormatClaude:
|
|
convertResult, err := relayconvert.ConvertResponse(c, info, types.RelayFormatClaude, fullTextResponse)
|
|
if err != nil {
|
|
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
|
|
}
|
|
claudeRespStr, err := common.Marshal(convertResult.Value)
|
|
if err != nil {
|
|
return nil, types.NewError(err, types.ErrorCodeBadResponseBody)
|
|
}
|
|
responseBody = claudeRespStr
|
|
case types.RelayFormatGemini:
|
|
break
|
|
}
|
|
|
|
service.IOCopyBytesGracefully(c, resp, responseBody)
|
|
|
|
return &usage, nil
|
|
}
|
|
|
|
func GeminiEmbeddingHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
|
|
defer service.CloseResponseBodyGracefully(resp)
|
|
|
|
responseBody, readErr := io.ReadAll(resp.Body)
|
|
if readErr != nil {
|
|
return nil, types.NewOpenAIError(readErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
|
|
var geminiResponse dto.GeminiBatchEmbeddingResponse
|
|
if jsonErr := common.Unmarshal(responseBody, &geminiResponse); jsonErr != nil {
|
|
return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
|
|
// convert to openai format response
|
|
openAIResponse := dto.OpenAIEmbeddingResponse{
|
|
Object: "list",
|
|
Data: make([]dto.OpenAIEmbeddingResponseItem, 0, len(geminiResponse.Embeddings)),
|
|
Model: info.UpstreamModelName,
|
|
}
|
|
|
|
for i, embedding := range geminiResponse.Embeddings {
|
|
openAIResponse.Data = append(openAIResponse.Data, dto.OpenAIEmbeddingResponseItem{
|
|
Object: "embedding",
|
|
Embedding: embedding.Values,
|
|
Index: i,
|
|
})
|
|
}
|
|
|
|
// calculate usage
|
|
// https://ai.google.dev/gemini-api/docs/pricing?hl=zh-cn#text-embedding-004
|
|
// Google has not yet clarified how embedding models will be billed
|
|
// refer to openai billing method to use input tokens billing
|
|
// https://platform.openai.com/docs/guides/embeddings#what-are-embeddings
|
|
usage := service.ResponseText2Usage(c, "", info.UpstreamModelName, info.GetEstimatePromptTokens())
|
|
openAIResponse.Usage = *usage
|
|
|
|
jsonResponse, jsonErr := common.Marshal(openAIResponse)
|
|
if jsonErr != nil {
|
|
return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
|
|
service.IOCopyBytesGracefully(c, resp, jsonResponse)
|
|
return usage, nil
|
|
}
|
|
|
|
func GeminiImageHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
|
|
responseBody, readErr := io.ReadAll(resp.Body)
|
|
if readErr != nil {
|
|
return nil, types.NewOpenAIError(readErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
_ = resp.Body.Close()
|
|
|
|
var geminiResponse dto.GeminiImageResponse
|
|
if jsonErr := common.Unmarshal(responseBody, &geminiResponse); jsonErr != nil {
|
|
return nil, types.NewOpenAIError(jsonErr, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
|
|
if len(geminiResponse.Predictions) == 0 {
|
|
return nil, types.NewOpenAIError(errors.New("no images generated"), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
|
|
}
|
|
|
|
// convert to openai format response
|
|
openAIResponse := dto.ImageResponse{
|
|
Created: common.GetTimestamp(),
|
|
Data: make([]dto.ImageData, 0, len(geminiResponse.Predictions)),
|
|
}
|
|
|
|
for _, prediction := range geminiResponse.Predictions {
|
|
if prediction.RaiFilteredReason != "" {
|
|
continue // skip filtered image
|
|
}
|
|
openAIResponse.Data = append(openAIResponse.Data, dto.ImageData{
|
|
B64Json: prediction.BytesBase64Encoded,
|
|
})
|
|
}
|
|
|
|
jsonResponse, jsonErr := common.Marshal(openAIResponse)
|
|
if jsonErr != nil {
|
|
return nil, types.NewError(jsonErr, types.ErrorCodeBadResponseBody)
|
|
}
|
|
|
|
c.Writer.Header().Set("Content-Type", "application/json")
|
|
c.Writer.WriteHeader(resp.StatusCode)
|
|
_, _ = c.Writer.Write(jsonResponse)
|
|
|
|
// https://github.com/google-gemini/cookbook/blob/719a27d752aac33f39de18a8d3cb42a70874917e/quickstarts/Counting_Tokens.ipynb
|
|
// each image has fixed 258 tokens
|
|
const imageTokens = 258
|
|
generatedImages := len(openAIResponse.Data)
|
|
|
|
usage := &dto.Usage{
|
|
PromptTokens: imageTokens * generatedImages, // each generated image has fixed 258 tokens
|
|
CompletionTokens: 0, // image generation does not calculate completion tokens
|
|
TotalTokens: imageTokens * generatedImages,
|
|
}
|
|
|
|
return usage, nil
|
|
}
|
|
|
|
type GeminiModelsResponse struct {
|
|
Models []dto.GeminiModel `json:"models"`
|
|
NextPageToken string `json:"nextPageToken"`
|
|
}
|
|
|
|
func FetchGeminiModels(baseURL, apiKey, proxyURL string) ([]string, error) {
|
|
client, err := service.GetHttpClientWithProxy(proxyURL)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("创建HTTP客户端失败: %v", err)
|
|
}
|
|
|
|
allModels := make([]string, 0)
|
|
nextPageToken := ""
|
|
maxPages := 100 // Safety limit to prevent infinite loops
|
|
|
|
for page := 0; page < maxPages; page++ {
|
|
url := fmt.Sprintf("%s/v1beta/models", baseURL)
|
|
if nextPageToken != "" {
|
|
url = fmt.Sprintf("%s?pageToken=%s", url, nextPageToken)
|
|
}
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
|
request, err := http.NewRequestWithContext(ctx, "GET", url, nil)
|
|
if err != nil {
|
|
cancel()
|
|
return nil, fmt.Errorf("创建请求失败: %v", err)
|
|
}
|
|
|
|
request.Header.Set("x-goog-api-key", apiKey)
|
|
|
|
response, err := client.Do(request)
|
|
if err != nil {
|
|
cancel()
|
|
return nil, fmt.Errorf("请求失败: %v", err)
|
|
}
|
|
|
|
if response.StatusCode != http.StatusOK {
|
|
body, _ := io.ReadAll(response.Body)
|
|
response.Body.Close()
|
|
cancel()
|
|
return nil, fmt.Errorf("服务器返回错误 %d: %s", response.StatusCode, string(body))
|
|
}
|
|
|
|
body, err := io.ReadAll(response.Body)
|
|
response.Body.Close()
|
|
cancel()
|
|
if err != nil {
|
|
return nil, fmt.Errorf("读取响应失败: %v", err)
|
|
}
|
|
|
|
var modelsResponse GeminiModelsResponse
|
|
if err = common.Unmarshal(body, &modelsResponse); err != nil {
|
|
return nil, fmt.Errorf("解析响应失败: %v", err)
|
|
}
|
|
|
|
for _, model := range modelsResponse.Models {
|
|
modelNameValue, ok := model.Name.(string)
|
|
if !ok {
|
|
continue
|
|
}
|
|
modelName := strings.TrimPrefix(modelNameValue, "models/")
|
|
allModels = append(allModels, modelName)
|
|
}
|
|
|
|
nextPageToken = modelsResponse.NextPageToken
|
|
if nextPageToken == "" {
|
|
break
|
|
}
|
|
}
|
|
|
|
return allModels, nil
|
|
}
|