Files
easyai-ai-gateway/apps/api/internal/httpapi/compat_protocol.go
T
easyai e07a997aa9 feat(api): 统一官方兼容接口响应协议
兼容接口现在以入口协议作为最终响应协议,同协议保留官方 Wire 响应,跨协议统一转换成功、任务状态与错误结构。

同时修正异步提交状态边界,持久化兼容公开任务标识和官方提交响应,并新增迁移、流式响应及协议契约测试。

验证:go vet ./...;go test ./...;govulncheck ./...;pnpm lint;pnpm test;pnpm build;pnpm audit --audit-level high;pnpm openapi;全部 CI 脚本。
2026-07-22 15:34:59 +08:00

181 lines
5.7 KiB
Go

package httpapi
import (
"encoding/json"
"errors"
"net/http"
"strings"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/clients"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
)
func (s *Server) requireProtocolUser(protocol string, next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
user, err := s.auth.Authenticate(r)
if err != nil {
status, code, message := http.StatusUnauthorized, "unauthorized", "unauthorized"
var requestErr *auth.RequestAuthError
if errors.As(err, &requestErr) {
status, code, message = requestErr.Status, requestErr.Code, requestErr.Message
}
writeProtocolError(w, protocol, status, message, nil, code)
return
}
if auth.PermissionLevel(user.Roles) < 1 {
writeProtocolError(w, protocol, http.StatusForbidden, "forbidden", nil, "permission_denied")
return
}
if strings.EqualFold(strings.TrimSpace(user.Source), "oidc") {
result, resolveErr := s.resolveOIDCUserProjection(r.Context(), r, user)
if resolveErr != nil {
status, code, message := http.StatusServiceUnavailable, errorCodeGatewayProvisioningFailed, "Gateway 账号初始化失败,请稍后重试"
switch {
case errors.Is(resolveErr, store.ErrOIDCUserNotProvisioned):
status, code, message = http.StatusForbidden, errorCodeGatewayUserNotProvisioned, "该账号尚未开通 EasyAI Gateway"
case errors.Is(resolveErr, store.ErrOIDCUserDisabled):
status, code, message = http.StatusForbidden, errorCodeGatewayUserDisabled, "该 Gateway 账号已停用,请联系管理员"
case errors.Is(resolveErr, store.ErrOIDCTenantUnavailable):
status, code, message = http.StatusServiceUnavailable, errorCodeGatewayTenantUnavailable, "Gateway 租户尚未就绪,请联系管理员"
}
writeProtocolError(w, protocol, status, message, nil, code)
return
}
user = result.User
}
next.ServeHTTP(w, r.WithContext(auth.WithUser(r.Context(), user)))
})
}
func targetProtocolForTaskRequest(kind string, r *http.Request) string {
switch kind {
case "chat.completions":
return clients.ProtocolOpenAIChatCompletions
case "responses":
return clients.ProtocolOpenAIResponses
case "embeddings":
return clients.ProtocolOpenAIEmbeddings
case "images.generations", "images.edits":
return clients.ProtocolOpenAIImages
default:
return ""
}
}
func wireResponseMatches(wire *clients.WireResponse, targetProtocol string) bool {
return wire != nil && !wire.Converted && strings.TrimSpace(wire.Protocol) == strings.TrimSpace(targetProtocol)
}
func writeWireResponse(w http.ResponseWriter, wire *clients.WireResponse) {
if wire == nil {
return
}
for name, values := range wire.Headers {
for _, value := range values {
if strings.EqualFold(name, "Content-Type") {
w.Header().Set(name, value)
} else {
w.Header().Add(name, value)
}
}
}
status := wire.StatusCode
if status == 0 {
status = http.StatusOK
}
w.WriteHeader(status)
if len(wire.RawJSON) > 0 {
_, _ = w.Write(wire.RawJSON)
return
}
_ = json.NewEncoder(w).Encode(wire.Body)
}
func writeProtocolError(w http.ResponseWriter, protocol string, status int, message string, details map[string]any, code string) {
switch protocol {
case clients.ProtocolGeminiGenerateContent:
writeGeminiError(w, status, message, details, code)
case clients.ProtocolVolcesContents:
writeVolcesError(w, status, message, code)
case clients.ProtocolKlingV1Omni, clients.ProtocolKlingV2Omni:
writeKelingCompatError(w, "", newKelingCompatError(status, kelingCompatBusinessCode(code, message), message))
default:
writeOpenAIError(w, status, message, details, code)
}
}
func writeOpenAIError(w http.ResponseWriter, status int, message string, details map[string]any, code string) {
errorType := "server_error"
if status >= 400 && status < 500 {
errorType = "invalid_request_error"
}
var param any
if details != nil {
param = details["param"]
}
payload := map[string]any{
"message": message,
"type": errorType,
"param": param,
"code": nil,
}
if strings.TrimSpace(code) != "" {
payload["code"] = code
}
writeJSON(w, status, map[string]any{"error": payload})
}
func writeGeminiError(w http.ResponseWriter, status int, message string, details map[string]any, code string) {
writeJSON(w, status, geminiErrorEnvelope(status, message, details, code))
}
func geminiErrorEnvelope(status int, message string, details map[string]any, code string) map[string]any {
detailList := []any{}
if len(details) > 0 {
detailList = append(detailList, details)
}
return map[string]any{"error": map[string]any{
"code": status,
"message": message,
"status": googleRPCStatus(status, code),
"details": detailList,
}}
}
func googleRPCStatus(status int, code string) string {
switch status {
case http.StatusBadRequest:
return "INVALID_ARGUMENT"
case http.StatusUnauthorized:
return "UNAUTHENTICATED"
case http.StatusForbidden:
return "PERMISSION_DENIED"
case http.StatusNotFound:
return "NOT_FOUND"
case http.StatusConflict:
return "ALREADY_EXISTS"
case http.StatusTooManyRequests:
return "RESOURCE_EXHAUSTED"
case http.StatusServiceUnavailable:
return "UNAVAILABLE"
case http.StatusGatewayTimeout:
return "DEADLINE_EXCEEDED"
default:
if strings.EqualFold(code, "cancelled") || strings.EqualFold(code, "canceled") {
return "CANCELLED"
}
return "INTERNAL"
}
}
func writeVolcesError(w http.ResponseWriter, status int, message string, code string) {
if strings.TrimSpace(code) == "" {
code = http.StatusText(status)
}
writeJSON(w, status, map[string]any{"error": map[string]any{
"code": code,
"message": message,
}})
}