78 lines
2.2 KiB
Go
78 lines
2.2 KiB
Go
package gemini
|
||
|
||
import (
|
||
"encoding/json"
|
||
"io"
|
||
"net/http"
|
||
"one-api/common"
|
||
"one-api/dto"
|
||
relaycommon "one-api/relay/common"
|
||
"one-api/service"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
)
|
||
|
||
func GeminiTextGenerationHandler(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (*dto.OpenAIErrorWithStatusCode, *dto.Usage) {
|
||
// 读取响应体
|
||
responseBody, err := io.ReadAll(resp.Body)
|
||
if err != nil {
|
||
return service.OpenAIErrorWrapper(err, "read_response_body_failed", http.StatusInternalServerError), nil
|
||
}
|
||
err = resp.Body.Close()
|
||
if err != nil {
|
||
return service.OpenAIErrorWrapper(err, "close_response_body_failed", http.StatusInternalServerError), nil
|
||
}
|
||
|
||
if common.DebugEnabled {
|
||
println(string(responseBody))
|
||
}
|
||
|
||
// 解析为 Gemini 原生响应格式
|
||
var geminiResponse dto.GeminiTextGenerationResponse
|
||
err = common.DecodeJson(responseBody, &geminiResponse)
|
||
if err != nil {
|
||
return service.OpenAIErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError), nil
|
||
}
|
||
|
||
// 检查是否有候选响应
|
||
if len(geminiResponse.Candidates) == 0 {
|
||
return &dto.OpenAIErrorWithStatusCode{
|
||
Error: dto.OpenAIError{
|
||
Message: "No candidates returned",
|
||
Type: "server_error",
|
||
Param: "",
|
||
Code: 500,
|
||
},
|
||
StatusCode: resp.StatusCode,
|
||
}, nil
|
||
}
|
||
|
||
// 计算使用量(基于 UsageMetadata)
|
||
usage := dto.Usage{
|
||
PromptTokens: geminiResponse.UsageMetadata.PromptTokenCount,
|
||
CompletionTokens: geminiResponse.UsageMetadata.CandidatesTokenCount,
|
||
TotalTokens: geminiResponse.UsageMetadata.TotalTokenCount,
|
||
}
|
||
|
||
// 设置模型版本
|
||
if geminiResponse.ModelVersion == "" {
|
||
geminiResponse.ModelVersion = info.UpstreamModelName
|
||
}
|
||
|
||
// 直接返回 Gemini 原生格式的 JSON 响应
|
||
jsonResponse, err := json.Marshal(geminiResponse)
|
||
if err != nil {
|
||
return service.OpenAIErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError), nil
|
||
}
|
||
|
||
// 设置响应头并写入响应
|
||
c.Writer.Header().Set("Content-Type", "application/json")
|
||
c.Writer.WriteHeader(resp.StatusCode)
|
||
_, err = c.Writer.Write(jsonResponse)
|
||
if err != nil {
|
||
return service.OpenAIErrorWrapper(err, "write_response_failed", http.StatusInternalServerError), nil
|
||
}
|
||
|
||
return nil, &usage
|
||
}
|