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:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user