feat(api): 统一官方兼容接口响应协议

兼容接口现在以入口协议作为最终响应协议,同协议保留官方 Wire 响应,跨协议统一转换成功、任务状态与错误结构。

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

验证:go vet ./...;go test ./...;govulncheck ./...;pnpm lint;pnpm test;pnpm build;pnpm audit --audit-level high;pnpm openapi;全部 CI 脚本。
This commit is contained in:
2026-07-22 15:34:59 +08:00
parent 42e8b517fd
commit e07a997aa9
32 changed files with 2436 additions and 591 deletions
+129 -13
View File
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"log/slog"
"net/http"
"os"
"regexp"
"strconv"
@@ -39,6 +40,7 @@ type billingMetricsObserver interface {
type Result struct {
Task store.GatewayTask
Output map[string]any
Wire *clients.WireResponse
}
var ErrTaskQueued = errors.New("task queued")
@@ -56,6 +58,10 @@ func (e *upstreamSubmissionUnknownError) Error() string {
return "upstream submission result is unknown"
}
func (e *upstreamSubmissionUnknownError) ErrorCode() string {
return "upstream_submission_unknown"
}
func (e *upstreamSubmissionUnknownError) Unwrap() error {
return e.Cause
}
@@ -151,22 +157,23 @@ func (s *Service) executeWithToken(ctx context.Context, task store.GatewayTask,
}
}
if err := validateRequest(task.Kind, body); err != nil {
validationErr := parameterPreprocessClientError(err)
s.recordFailedAttempt(ctx, failedAttemptRecord{
Task: task,
Body: body,
AttemptNo: task.AttemptCount + 1,
Code: "bad_request",
Cause: err,
Cause: validationErr,
Simulated: task.RunMode == "simulation",
Scope: "request_validation",
Reason: "request_validation_failed",
ModelType: modelType,
})
failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "bad_request", err.Error(), task.RunMode == "simulation", err)
failed, finishErr := s.failTask(ctx, task.ID, task.ExecutionToken, "bad_request", validationErr.Error(), task.RunMode == "simulation", validationErr)
if finishErr != nil {
return Result{}, finishErr
}
return Result{Task: failed, Output: failed.Result}, err
return Result{Task: failed, Output: failed.Result}, validationErr
}
var clonedVoice clonedVoiceBinding
body, clonedVoice, err = s.resolveClonedVoiceBinding(ctx, user, task.Kind, body)
@@ -579,7 +586,7 @@ candidatesLoop:
}
walletReservationFinalized = true
s.logger.Warn("task succeeded but billing requires manual review", "taskID", task.ID, "error_category", "billing_calculation_failed")
return Result{Task: review, Output: response.Result}, nil
return Result{Task: review, Output: response.Result, Wire: response.Wire}, nil
}
finalAmountText = finalAmount.String()
} else {
@@ -637,7 +644,7 @@ candidatesLoop:
}, isSimulation(task, candidate)); err != nil {
return Result{}, err
}
return Result{Task: finished, Output: response.Result}, nil
return Result{Task: finished, Output: response.Result, Wire: response.Wire}, nil
}
var submissionUnknown *upstreamSubmissionUnknownError
if errors.As(err, &submissionUnknown) {
@@ -986,9 +993,27 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
publicResponseID = responseExecution.PublicResponseID
publicPreviousResponseID = responseExecution.PublicPreviousResponseID
}
if err := s.store.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, "submitting"); err != nil {
return clients.Response{}, fmt.Errorf("mark upstream submission: %w", err)
submissionStatus := "not_submitted"
if strings.TrimSpace(task.RemoteTaskID) != "" {
submissionStatus = "response_received"
if err := s.store.SetAttemptUpstreamSubmissionStatus(ctx, attemptID, submissionStatus); err != nil {
return clients.Response{}, fmt.Errorf("restore upstream submission status: %w", err)
}
}
setSubmissionStatus := func(status string) error {
if submissionStatus == "response_received" && status != "response_received" {
return nil
}
if submissionStatus == status {
return nil
}
if err := s.store.SetAttemptUpstreamSubmissionStatus(context.WithoutCancel(ctx), attemptID, status); err != nil {
return err
}
submissionStatus = status
return nil
}
var submissionWire *clients.WireResponse
response, err := client.Run(ctx, clients.Request{
Kind: task.Kind,
ModelType: candidate.ModelType,
@@ -1002,13 +1027,36 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
if strings.TrimSpace(remoteTaskID) == "" {
return nil
}
return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload)
if err := s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload); err != nil {
return err
}
task.RemoteTaskID = remoteTaskID
task.RemoteTaskPayload = payload
if err := s.persistCompatibilitySubmission(context.WithoutCancel(ctx), task, candidate, remoteTaskID, payload, submissionWire); err != nil {
return err
}
return setSubmissionStatus("response_received")
},
OnRemoteTaskPolled: func(remoteTaskID string, payload map[string]any) error {
if strings.TrimSpace(remoteTaskID) == "" {
return nil
}
return s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload)
if err := s.store.SetTaskRemoteTask(context.WithoutCancel(ctx), task.ID, task.ExecutionToken, attemptID, remoteTaskID, payload); err != nil {
return err
}
task.RemoteTaskID = remoteTaskID
task.RemoteTaskPayload = payload
return setSubmissionStatus("response_received")
},
OnUpstreamSubmissionStarted: func() error {
return setSubmissionStatus("submitting")
},
OnUpstreamResponseReceived: func() error {
return setSubmissionStatus("response_received")
},
OnUpstreamWireResponse: func(wire *clients.WireResponse) error {
submissionWire = wire
return s.persistCompatibilitySubmission(context.WithoutCancel(ctx), task, candidate, "", wire.Body, wire)
},
Stream: boolFromMap(providerBody, "stream"),
StreamDelta: onDelta,
@@ -1020,7 +1068,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
})
callFinishedAt := time.Now()
if err == nil {
if markErr := s.store.SetAttemptUpstreamSubmissionStatus(context.WithoutCancel(ctx), attemptID, "response_received"); markErr != nil {
if markErr := setSubmissionStatus("response_received"); markErr != nil {
return clients.Response{}, &upstreamSubmissionUnknownError{AttemptID: attemptID, Cause: markErr}
}
}
@@ -1037,8 +1085,8 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
}
}
if err != nil {
if clients.ErrorResponseMetadata(err).StatusCode > 0 {
if markErr := s.store.SetAttemptUpstreamSubmissionStatus(context.WithoutCancel(ctx), attemptID, "response_received"); markErr != nil {
if clients.ErrorResponseMetadata(err).StatusCode > 0 && submissionStatus != "response_received" {
if markErr := setSubmissionStatus("response_received"); markErr != nil {
return clients.Response{}, &upstreamSubmissionUnknownError{AttemptID: attemptID, Cause: markErr}
}
}
@@ -1070,7 +1118,7 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
ErrorMessage: err.Error(),
})
_ = s.emit(ctx, task.ID, "task.attempt.failed", "running", "attempt_failed", 0.45, err.Error(), map[string]any{"attempt": attemptNo, "retryable": retryable, "requestId": requestID, "statusCode": clients.ErrorResponseMetadata(err).StatusCode, "metrics": metrics}, simulated)
if !simulated && clients.ErrorResponseMetadata(err).StatusCode == 0 {
if !simulated && submissionStatus == "submitting" {
return clients.Response{}, &upstreamSubmissionUnknownError{AttemptID: attemptID, Cause: err}
}
s.applyCandidateFailurePolicies(ctx, task.ID, candidate, err, simulated, singleSourceProtected)
@@ -1172,6 +1220,74 @@ func (s *Service) runCandidate(ctx context.Context, task store.GatewayTask, user
return response, nil
}
func (s *Service) persistCompatibilitySubmission(ctx context.Context, task store.GatewayTask, candidate store.RuntimeModelCandidate, remoteTaskID string, payload map[string]any, wire *clients.WireResponse) error {
targetProtocol := strings.TrimSpace(stringFromMap(task.Request, "_gateway_target_protocol"))
if targetProtocol == "" {
return nil
}
sourceProtocol := compatibilitySourceProtocol(task.Kind, candidate)
if wire != nil && strings.TrimSpace(wire.Protocol) != "" {
sourceProtocol = strings.TrimSpace(wire.Protocol)
}
publicID := task.ID
if sourceProtocol == targetProtocol && strings.TrimSpace(remoteTaskID) != "" {
publicID = strings.TrimSpace(remoteTaskID)
}
submitBody := payload
if nested, ok := payload["submit"].(map[string]any); ok && len(nested) > 0 {
submitBody = nested
}
httpStatus := http.StatusOK
headers := map[string]any{}
if wire != nil {
httpStatus = wire.StatusCode
if len(wire.Body) > 0 {
submitBody = wire.Body
}
for name, values := range wire.Headers {
headers[name] = append([]string(nil), values...)
}
}
return s.store.SetTaskCompatibilitySubmission(ctx, task.ID, store.CompatibilitySubmission{
TargetProtocol: targetProtocol,
PublicID: publicID,
SourceProtocol: sourceProtocol,
HTTPStatus: httpStatus,
Headers: headers,
Body: submitBody,
})
}
func compatibilitySourceProtocol(kind string, candidate store.RuntimeModelCandidate) string {
if protocol := strings.TrimSpace(candidate.ResponseProtocol); protocol != "" {
return protocol
}
switch strings.ToLower(strings.TrimSpace(candidate.Provider)) {
case "volces":
if kind == "videos.generations" {
return clients.ProtocolVolcesContents
}
case "keling", "kling":
if kind == "videos.generations" {
return clients.ProtocolKlingV1Omni
}
case "gemini":
return clients.ProtocolGeminiGenerateContent
case "openai":
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
}
}
return strings.ToLower(strings.TrimSpace(candidate.Provider))
}
func (s *Service) recordTaskParameterPreprocessing(ctx context.Context, taskID string, attemptID string, attemptNo int, candidate store.RuntimeModelCandidate, log parameterPreprocessingLog) error {
if skipTaskParameterPreprocessingLog(log.ModelType) && !log.Changed {
return nil