diff --git a/internal/core/model/models.go b/internal/core/model/models.go index eb31630..329ec89 100644 --- a/internal/core/model/models.go +++ b/internal/core/model/models.go @@ -129,7 +129,11 @@ type GenerationOutput struct { func (GenerationOutput) TableName() string { return "generation_outputs" } type Attempt struct { + Type string `json:"type,omitempty"` + RouteMemberID uint64 `json:"route_member_id,omitempty"` ProviderModelID uint64 `json:"provider_model_id,omitempty"` + ProviderOrdinal uint32 `json:"provider_ordinal,omitempty"` + Retryable *bool `json:"retryable,omitempty"` ErrorCode string `json:"error_code,omitempty"` ErrorMessage string `json:"error_message,omitempty"` LatencyMS int64 `json:"latency_ms,omitempty"` diff --git a/internal/core/provider/openai.go b/internal/core/provider/openai.go index 9ab50f1..009599d 100644 --- a/internal/core/provider/openai.go +++ b/internal/core/provider/openai.go @@ -12,6 +12,8 @@ import ( "net/url" "path" "strings" + + "git.ilapage.cn/OPC/chorus/internal/core/model" ) const defaultMaxResponseBytes int64 = 20 << 20 @@ -21,8 +23,11 @@ var ( ErrInvalidRequest = errors.New("provider request is invalid") ) +// OpenAIConfig applies to both OpenAI-compatible and Gemini protocol shapes. +// The name is retained for MVP-0 callers; APIType controls the fixed endpoint. type OpenAIConfig struct { BaseURL string + AuthType AuthType APIKey string ExtraBody json.RawMessage MaxResponseBytes int64 @@ -31,6 +36,7 @@ type OpenAIConfig struct { type OpenAI struct { http HTTPClient baseURL *url.URL + authType AuthType apiKey string extraBody map[string]any maxResponseBytes int64 @@ -44,6 +50,17 @@ func NewOpenAI(httpClient HTTPClient, config OpenAIConfig) (*OpenAI, error) { if err != nil || (baseURL.Scheme != "http" && baseURL.Scheme != "https") || baseURL.Hostname() == "" || baseURL.User != nil || baseURL.RawQuery != "" || baseURL.Fragment != "" { return nil, ErrInvalidConfig } + authType := config.AuthType + if authType == "" { + if config.APIKey == "" { + authType = AuthNone + } else { + authType = AuthBearer + } + } + if !authType.Valid() || (authType != AuthNone && strings.TrimSpace(config.APIKey) == "") { + return nil, ErrInvalidConfig + } extraBody, err := parseExtraBody(config.ExtraBody) if err != nil { return nil, err @@ -51,7 +68,7 @@ func NewOpenAI(httpClient HTTPClient, config OpenAIConfig) (*OpenAI, error) { if config.MaxResponseBytes <= 0 { config.MaxResponseBytes = defaultMaxResponseBytes } - return &OpenAI{http: httpClient, baseURL: baseURL, apiKey: config.APIKey, extraBody: extraBody, maxResponseBytes: config.MaxResponseBytes}, nil + return &OpenAI{http: httpClient, baseURL: baseURL, authType: authType, apiKey: config.APIKey, extraBody: extraBody, maxResponseBytes: config.MaxResponseBytes}, nil } func (c *OpenAI) Generate(ctx context.Context, request Request) ([]Output, error) { @@ -59,10 +76,29 @@ func (c *OpenAI) Generate(ctx context.Context, request Request) ([]Output, error return nil, ErrInvalidRequest } switch request.APIType { - case "chat": + case model.APIChat: + if request.Kind != model.KindText || len(request.Inputs) != 0 { + return nil, ErrInvalidRequest + } return c.chat(ctx, request) - case "images_edits": + case model.APIImages: + if request.Kind != model.KindImage || len(request.Inputs) != 0 { + return nil, ErrInvalidRequest + } + return c.images(ctx, request) + case model.APIImagesEdits: + if request.Kind != model.KindImage || !validImageInputs(request.Inputs, true) { + return nil, ErrInvalidRequest + } return c.imagesEdits(ctx, request) + case model.APIGemini: + if request.Kind != model.KindText && request.Kind != model.KindImage { + return nil, ErrInvalidRequest + } + if len(request.Inputs) > 0 && !validImageInputs(request.Inputs, true) { + return nil, ErrInvalidRequest + } + return c.gemini(ctx, request) default: return nil, ErrInvalidRequest } @@ -78,7 +114,7 @@ func (c *OpenAI) chat(ctx context.Context, request Request) ([]Output, error) { if err != nil { return nil, ErrInvalidRequest } - response, err := c.do(ctx, "/chat/completions", "application/json", encoded) + response, err := c.do(ctx, c.openAIURL("/chat/completions"), "application/json", encoded) if err != nil { return nil, err } @@ -95,10 +131,26 @@ func (c *OpenAI) chat(ctx context.Context, request Request) ([]Output, error) { return []Output{{Kind: request.Kind, Text: decoded.Choices[0].Message.Content, ContentType: "text/plain; charset=utf-8"}}, nil } -func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, error) { - if len(request.Inputs) == 0 { +func (c *OpenAI) images(ctx context.Context, request Request) ([]Output, error) { + body := map[string]any{ + "model": request.ModelID, + "prompt": request.RenderedPrompt, + "n": 1, + "response_format": "b64_json", + } + mergeExtra(body, c.extraBody) + encoded, err := json.Marshal(body) + if err != nil { return nil, ErrInvalidRequest } + response, err := c.do(ctx, c.openAIURL("/images/generations"), "application/json", encoded) + if err != nil { + return nil, err + } + return c.decodeOpenAIImages(ctx, request.Kind, response, true) +} + +func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, error) { var body bytes.Buffer writer := multipart.NewWriter(&body) if err := writer.WriteField("model", request.ModelID); err != nil { @@ -114,9 +166,6 @@ func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, er } } for index, input := range request.Inputs { - if len(input.Content) == 0 || !strings.HasPrefix(input.MIMEType, "image/") { - return nil, ErrInvalidRequest - } part, err := writer.CreateFormFile("image[]", fmt.Sprintf("image-%d%s", index+1, imageExtension(input.MIMEType))) if err != nil { return nil, ErrInvalidRequest @@ -128,29 +177,90 @@ func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, er if err := writer.Close(); err != nil { return nil, ErrInvalidRequest } - response, err := c.do(ctx, "/images/edits", writer.FormDataContentType(), body.Bytes()) + response, err := c.do(ctx, c.openAIURL("/images/edits"), writer.FormDataContentType(), body.Bytes()) if err != nil { return nil, err } + return c.decodeOpenAIImages(ctx, request.Kind, response, false) +} + +func (c *OpenAI) gemini(ctx context.Context, request Request) ([]Output, error) { + parts := make([]map[string]any, 0, len(request.Inputs)+1) + parts = append(parts, map[string]any{"text": request.RenderedPrompt}) + for _, input := range request.Inputs { + parts = append(parts, map[string]any{"inlineData": map[string]string{ + "mimeType": input.MIMEType, + "data": base64.StdEncoding.EncodeToString(input.Content), + }}) + } + body := map[string]any{"contents": []any{map[string]any{"role": "user", "parts": parts}}} + mergeExtra(body, c.extraBody) + encoded, err := json.Marshal(body) + if err != nil { + return nil, ErrInvalidRequest + } + response, err := c.do(ctx, c.geminiURL(request.ModelID), "application/json", encoded) + if err != nil { + return nil, err + } + var decoded struct { + Candidates []struct { + Content struct { + Parts []struct { + Text string `json:"text"` + InlineData *struct { + MIMEType string `json:"mimeType"` + Data string `json:"data"` + } `json:"inlineData"` + } `json:"parts"` + } `json:"content"` + } `json:"candidates"` + } + if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Candidates) == 0 { + return nil, providerError(CodeUnknown, FailureOther) + } + for _, candidate := range decoded.Candidates { + for _, part := range candidate.Content.Parts { + if request.Kind == model.KindText && part.Text != "" { + return []Output{{Kind: request.Kind, Text: part.Text, ContentType: "text/plain; charset=utf-8"}}, nil + } + if request.Kind == model.KindImage && part.InlineData != nil { + if !strings.HasPrefix(strings.ToLower(part.InlineData.MIMEType), "image/") { + return nil, providerError(CodeUnknown, FailureOther) + } + content, decodeErr := base64.StdEncoding.DecodeString(part.InlineData.Data) + if decodeErr != nil || int64(len(content)) > c.maxResponseBytes { + return nil, providerError(CodeUnknown, FailureOther) + } + return []Output{{Kind: request.Kind, Content: content, ContentType: part.InlineData.MIMEType}}, nil + } + } + } + return nil, providerError(CodeUnknown, FailureOther) +} + +func (c *OpenAI) decodeOpenAIImages(ctx context.Context, kind model.GenerationKind, response HTTPResponse, requireOne bool) ([]Output, error) { var decoded struct { Data []struct { B64JSON string `json:"b64_json"` URL string `json:"url"` } `json:"data"` } - if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Data) == 0 { + if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Data) == 0 || (requireOne && len(decoded.Data) != 1) { return nil, providerError(CodeUnknown, FailureOther) } outputs := make([]Output, 0, len(decoded.Data)) for _, item := range decoded.Data { var content []byte contentType := "image/png" - if item.B64JSON != "" { + var err error + switch { + case item.B64JSON != "": content, err = base64.StdEncoding.DecodeString(item.B64JSON) - if err != nil { + if err != nil || int64(len(content)) > c.maxResponseBytes { return nil, providerError(CodeUnknown, FailureOther) } - } else if item.URL != "" { + case item.URL != "": fetched, fetchErr := c.http.Fetch(ctx, item.URL, c.maxResponseBytes) if fetchErr != nil { return nil, classifyNetworkError(fetchErr) @@ -159,22 +269,26 @@ func (c *OpenAI) imagesEdits(ctx context.Context, request Request) ([]Output, er return nil, fromHTTPStatus(fetched.StatusCode, nil) } content, contentType = fetched.Body, strings.TrimSpace(strings.Split(fetched.ContentType, ";")[0]) - } else { + if !strings.HasPrefix(strings.ToLower(contentType), "image/") || int64(len(content)) > c.maxResponseBytes { + return nil, providerError(CodeUnknown, FailureOther) + } + default: return nil, providerError(CodeUnknown, FailureOther) } - outputs = append(outputs, Output{Kind: request.Kind, Content: content, ContentType: contentType}) + outputs = append(outputs, Output{Kind: kind, Content: content, ContentType: contentType}) } return outputs, nil } -func (c *OpenAI) do(ctx context.Context, endpoint, contentType string, body []byte) (HTTPResponse, error) { - target := *c.baseURL - target.Path = path.Join(strings.TrimSuffix(c.baseURL.Path, "/"), endpoint) +func (c *OpenAI) do(ctx context.Context, target string, contentType string, body []byte) (HTTPResponse, error) { header := http.Header{"Content-Type": []string{contentType}, "Accept": []string{"application/json"}} - if c.apiKey != "" { + switch c.authType { + case AuthBearer: header.Set("Authorization", "Bearer "+c.apiKey) + case AuthGoogleAPIKey: + header.Set("x-goog-api-key", c.apiKey) } - response, err := c.http.Do(ctx, HTTPRequest{Method: http.MethodPost, URL: target.String(), Header: header, Body: body, MaxBytes: c.maxResponseBytes}) + response, err := c.http.Do(ctx, HTTPRequest{Method: http.MethodPost, URL: target, Header: header, Body: body, MaxBytes: c.maxResponseBytes}) if err != nil { return HTTPResponse{}, classifyNetworkError(err) } @@ -184,6 +298,43 @@ func (c *OpenAI) do(ctx context.Context, endpoint, contentType string, body []by return response, nil } +func (c *OpenAI) openAIURL(endpoint string) string { + target := *c.baseURL + target.Path = path.Join(strings.TrimSuffix(c.baseURL.Path, "/"), endpoint) + target.RawPath = "" + return target.String() +} + +func (c *OpenAI) geminiURL(modelID string) string { + target := *c.baseURL + target.Path = path.Join(strings.TrimSuffix(c.baseURL.Path, "/"), "models", modelID+":generateContent") + target.RawPath = path.Join(strings.TrimSuffix(c.baseURL.EscapedPath(), "/"), "models", url.PathEscape(modelID)+":generateContent") + return target.String() +} + +func validImageInputs(inputs []Input, requirePrimary bool) bool { + if len(inputs) == 0 { + return !requirePrimary + } + primary := 0 + positions := make(map[uint32]struct{}, len(inputs)) + for _, input := range inputs { + if len(input.Content) == 0 || !strings.HasPrefix(strings.ToLower(input.MIMEType), "image/") { + return false + } + if _, exists := positions[input.Position]; exists { + return false + } + positions[input.Position] = struct{}{} + if input.Role == model.RolePrimary { + primary++ + } else if input.Role != model.RoleReference { + return false + } + } + return !requirePrimary || primary == 1 +} + func parseExtraBody(raw json.RawMessage) (map[string]any, error) { result := map[string]any{} if len(raw) == 0 || string(raw) == "null" { @@ -192,9 +343,9 @@ func parseExtraBody(raw json.RawMessage) (map[string]any, error) { if err := json.Unmarshal(raw, &result); err != nil { return nil, ErrInvalidConfig } - allowed := map[string]bool{"temperature": true, "max_tokens": true, "size": true, "quality": true, "response_format": true, "n": true} + allowed := map[string]bool{"temperature": true, "max_tokens": true, "size": true, "quality": true} for key, value := range result { - if !allowed[key] || key == "model" || key == "messages" || key == "prompt" { + if !allowed[key] { return nil, ErrInvalidConfig } switch value.(type) { @@ -211,6 +362,7 @@ func mergeExtra(target, extra map[string]any) { target[key] = value } } + func scalarString(value any) (string, error) { switch value := value.(type) { case string: @@ -223,6 +375,7 @@ func scalarString(value any) (string, error) { return "", ErrInvalidConfig } } + func imageExtension(mimeType string) string { if mimeType == "image/jpeg" { return ".jpg" @@ -247,16 +400,19 @@ func fromHTTPStatus(status int, body []byte) *Error { } return providerError(CodeUnknown, FailureOther) } + func policyRejected(body []byte) bool { lower := strings.ToLower(string(body)) return strings.Contains(lower, "content_policy") || strings.Contains(lower, "safety") } + func classifyNetworkError(err error) *Error { if errors.Is(err, context.DeadlineExceeded) { return providerError(CodeTimeout, FailureTimeout) } return providerError(CodeConnection, FailureConnection) } + func providerError(code ErrorCode, class FailureClass) *Error { return &Error{Code: code, Class: class, Message: "upstream request failed"} } diff --git a/internal/core/provider/openai_test.go b/internal/core/provider/openai_test.go index 32154d4..27823c3 100644 --- a/internal/core/provider/openai_test.go +++ b/internal/core/provider/openai_test.go @@ -48,7 +48,7 @@ func TestImagesEditsBase64AndURL(t *testing.T) { responseBody := `{"data":[{"b64_json":"` + base64.StdEncoding.EncodeToString(pngData) + `"},{"url":"https://result.test/output"}]}` httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(responseBody)}, fetched: HTTPResponse{StatusCode: 200, ContentType: "image/png", Body: pngData}} client, _ := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: jsonBytes(`{"size":"1024x1024"}`)}) - outputs, err := client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImagesEdits, ModelID: "mock-image", RenderedPrompt: "edit", Inputs: []Input{{MIMEType: "image/png", Content: pngData}}}) + outputs, err := client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImagesEdits, ModelID: "mock-image", RenderedPrompt: "edit", Inputs: []Input{{Role: model.RolePrimary, MIMEType: "image/png", Content: pngData}}}) if err != nil || len(outputs) != 2 || string(outputs[0].Content) != "image-bytes" || httpClient.fetchURL != "https://result.test/output" { t.Fatalf("Generate() = %#v, %v, fetch=%q", outputs, err, httpClient.fetchURL) } diff --git a/internal/core/provider/protocol_test.go b/internal/core/provider/protocol_test.go new file mode 100644 index 0000000..93e9cdc --- /dev/null +++ b/internal/core/provider/protocol_test.go @@ -0,0 +1,97 @@ +package provider + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "strings" + "testing" + + "git.ilapage.cn/OPC/chorus/internal/core/model" +) + +func TestImagesProtocolForcesOneResultAndGoogleAuthentication(t *testing.T) { + image := base64.StdEncoding.EncodeToString([]byte("png")) + httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(`{"data":[{"b64_json":"` + image + `"}]}`)}} + client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", AuthType: AuthGoogleAPIKey, APIKey: "test-key"}) + if err != nil { + t.Fatal(err) + } + outputs, err := client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImages, ModelID: "mock-image", RenderedPrompt: "draw"}) + if err != nil || len(outputs) != 1 || string(outputs[0].Content) != "png" { + t.Fatalf("Generate() = %#v, %v", outputs, err) + } + if httpClient.request.URL != "https://provider.test/v1/images/generations" || httpClient.request.Header.Get("x-goog-api-key") != "test-key" || httpClient.request.Header.Get("Authorization") != "" { + t.Fatalf("request=%#v", httpClient.request) + } + var body map[string]any + if err := json.Unmarshal(httpClient.request.Body, &body); err != nil || body["n"] != float64(1) || body["response_format"] != "b64_json" { + t.Fatalf("fixed images request body=%s error=%v", httpClient.request.Body, err) + } +} + +func TestImagesRejectsMultipleResultsAndOversizedBase64(t *testing.T) { + encoded := base64.StdEncoding.EncodeToString([]byte("image")) + for _, test := range []struct { + name string + body string + limit int64 + }{ + {"multiple", `{"data":[{"b64_json":"` + encoded + `"},{"b64_json":"` + encoded + `"}]}`, 1024}, + {"too_large", `{"data":[{"b64_json":"` + encoded + `"}]}`, 2}, + } { + t.Run(test.name, func(t *testing.T) { + client, err := NewOpenAI(&fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(test.body)}}, OpenAIConfig{BaseURL: "https://provider.test/v1", MaxResponseBytes: test.limit}) + if err != nil { + t.Fatal(err) + } + _, err = client.Generate(context.Background(), Request{Kind: model.KindImage, APIType: model.APIImages, ModelID: "m", RenderedPrompt: "p"}) + var providerErr *Error + if !errors.As(err, &providerErr) || providerErr.Code != CodeUnknown { + t.Fatalf("error=%v", err) + } + }) + } +} + +func TestGeminiProtocolEscapesModelAndUsesInlineData(t *testing.T) { + image := []byte("image-content") + response := `{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"` + base64.StdEncoding.EncodeToString(image) + `"}}]}}]}` + httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, Body: []byte(response)}} + client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1beta", AuthType: AuthNone}) + if err != nil { + t.Fatal(err) + } + outputs, err := client.Generate(context.Background(), Request{ + Kind: model.KindImage, APIType: model.APIGemini, ModelID: "gemini/image", RenderedPrompt: "edit", + Inputs: []Input{{Role: model.RolePrimary, Position: 0, MIMEType: "image/png", Content: image}}, + }) + if err != nil || len(outputs) != 1 || string(outputs[0].Content) != string(image) { + t.Fatalf("Generate() = %#v, %v", outputs, err) + } + if httpClient.request.URL != "https://provider.test/v1beta/models/gemini%2Fimage:generateContent" || !strings.Contains(string(httpClient.request.Body), `"inlineData"`) || !strings.Contains(string(httpClient.request.Body), base64.StdEncoding.EncodeToString(image)) { + t.Fatalf("gemini request=%#v body=%s", httpClient.request, httpClient.request.Body) + } +} + +func TestProtocolCapabilityAndReservedBodyValidation(t *testing.T) { + client, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1"}) + if err != nil { + t.Fatal(err) + } + for _, request := range []Request{ + {Kind: model.KindImage, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"}, + {Kind: model.KindImage, APIType: model.APIImages, ModelID: "m", RenderedPrompt: "p", Inputs: []Input{{Role: model.RolePrimary, MIMEType: "image/png", Content: []byte("x")}}}, + {Kind: model.KindImage, APIType: model.APIImagesEdits, ModelID: "m", RenderedPrompt: "p", Inputs: []Input{{Role: model.RoleReference, MIMEType: "image/png", Content: []byte("x")}}}, + } { + if _, err := client.Generate(context.Background(), request); !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("request=%#v error=%v", request, err) + } + } + for _, value := range []string{`{"n":2}`, `{"response_format":"url"}`, `{"authorization":"override"}`} { + if _, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: []byte(value)}); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("extra body %s error=%v", value, err) + } + } +} diff --git a/internal/core/provider/provider.go b/internal/core/provider/provider.go index b0c7c95..943aad4 100644 --- a/internal/core/provider/provider.go +++ b/internal/core/provider/provider.go @@ -40,6 +40,20 @@ type Error struct { Message string } +// AuthType is deliberately small. Provider configuration must not supply +// arbitrary request headers because credentials are a security boundary. +type AuthType string + +const ( + AuthNone AuthType = "none" + AuthBearer AuthType = "bearer" + AuthGoogleAPIKey AuthType = "x-goog-api-key" +) + +func (a AuthType) Valid() bool { + return a == AuthNone || a == AuthBearer || a == AuthGoogleAPIKey +} + func (e *Error) Error() string { return fmt.Sprintf("provider request failed: %s", e.Code) } diff --git a/internal/core/queue/controller.go b/internal/core/queue/controller.go index 169812f..0313df8 100644 --- a/internal/core/queue/controller.go +++ b/internal/core/queue/controller.go @@ -40,6 +40,14 @@ func (c *Controller) AssignProvider(ctx context.Context, generationID uint64, le return c.repository.AssignProvider(ctx, generationID, leaseToken, providerModelID) } +func (c *Controller) BeginProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, routeMemberID, providerModelID uint64) (model.Attempt, bool, error) { + return c.repository.BeginProviderAttempt(ctx, generationID, leaseToken, routeMemberID, providerModelID) +} + +func (c *Controller) FinishProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, attempt model.Attempt) (bool, error) { + return c.repository.FinishProviderAttempt(ctx, generationID, leaseToken, attempt) +} + func (c *Controller) Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) { return c.repository.Succeed(ctx, generationID, leaseToken, outputs, attempt) } diff --git a/internal/core/queue/controller_test.go b/internal/core/queue/controller_test.go index 9c28c6d..86a5d12 100644 --- a/internal/core/queue/controller_test.go +++ b/internal/core/queue/controller_test.go @@ -19,6 +19,12 @@ func (f *fakeRepository) ClaimNext(context.Context, string, time.Duration) (*Cla f.claims++ return nil, nil } +func (f *fakeRepository) BeginProviderAttempt(context.Context, uint64, string, uint64, uint64) (model.Attempt, bool, error) { + return model.Attempt{}, false, nil +} +func (f *fakeRepository) FinishProviderAttempt(context.Context, uint64, string, model.Attempt) (bool, error) { + return false, nil +} func (f *fakeRepository) AssignProvider(context.Context, uint64, string, uint64) (bool, error) { return false, nil } diff --git a/internal/core/queue/mysql_integration_test.go b/internal/core/queue/mysql_integration_test.go index 5947391..d874d29 100644 --- a/internal/core/queue/mysql_integration_test.go +++ b/internal/core/queue/mysql_integration_test.go @@ -331,3 +331,50 @@ func TestMySQLFailureSanitizationAndCompletionRollback(t *testing.T) { t.Fatalf("valid Succeed() = %v, %v", owned, err) } } + +func TestMySQLProviderAttemptCASAndCounters(t *testing.T) { + repository, db := integrationRepository(t) + user := createIntegrationUser(t, db, "provider-attempt") + providerModel := createIntegrationProviderModel(t, db, "provider-attempt") + generation := createPending(t, repository, user.ID, "provider-attempt") + firstClaim, err := repository.ClaimNext(context.Background(), "worker-a", time.Second) + if err != nil || firstClaim == nil || firstClaim.Generation.ID != generation.ID { + t.Fatalf("first claim=%#v error=%v", firstClaim, err) + } + firstAttempt, owned, err := repository.BeginProviderAttempt(context.Background(), generation.ID, firstClaim.LeaseToken, 1, providerModel.ID) + if err != nil || !owned || firstAttempt.ProviderOrdinal != 1 || firstAttempt.Type != "provider" { + t.Fatalf("first begin=%#v owned=%v error=%v", firstAttempt, owned, err) + } + if err := db.Model(&model.Generation{}).Where("id = ?", generation.ID).Update("lease_until", time.Now().Add(-time.Second)).Error; err != nil { + t.Fatal(err) + } + secondClaim, err := repository.ClaimNext(context.Background(), "worker-b", time.Second) + if err != nil || secondClaim == nil || secondClaim.LeaseToken == firstClaim.LeaseToken { + t.Fatalf("second claim=%#v error=%v", secondClaim, err) + } + firstAttempt.ErrorCode = "upstream_timeout" + if owned, err := repository.FinishProviderAttempt(context.Background(), generation.ID, firstClaim.LeaseToken, firstAttempt); err != nil || owned { + t.Fatalf("stale finish owned=%v error=%v", owned, err) + } + secondAttempt, owned, err := repository.BeginProviderAttempt(context.Background(), generation.ID, secondClaim.LeaseToken, 1, providerModel.ID) + if err != nil || !owned || secondAttempt.ProviderOrdinal != 2 { + t.Fatalf("second begin=%#v owned=%v error=%v", secondAttempt, owned, err) + } + if owned, err := repository.FinishProviderAttempt(context.Background(), generation.ID, secondClaim.LeaseToken, secondAttempt); err != nil || !owned { + t.Fatalf("second finish owned=%v error=%v", owned, err) + } + if owned, err := repository.Succeed(context.Background(), generation.ID, secondClaim.LeaseToken, []model.GenerationOutput{textOutput("done")}, model.Attempt{}); err != nil || !owned { + t.Fatalf("succeed owned=%v error=%v", owned, err) + } + var stored model.Generation + if err := db.First(&stored, generation.ID).Error; err != nil { + t.Fatal(err) + } + if stored.Status != model.StatusSucceeded || stored.ProviderAttemptCount != 2 || stored.AttemptCount != 4 { + t.Fatalf("stored=%#v", stored) + } + var attempts []model.Attempt + if err := json.Unmarshal(stored.Attempts, &attempts); err != nil || len(attempts) != 4 || attempts[1].ErrorCode != "lease_expired" || attempts[3].FinishedAt == nil { + t.Fatalf("attempts=%#v error=%v", attempts, err) + } +} diff --git a/internal/core/queue/mysql_repository.go b/internal/core/queue/mysql_repository.go index 535ee90..a4ee952 100644 --- a/internal/core/queue/mysql_repository.go +++ b/internal/core/queue/mysql_repository.go @@ -122,7 +122,7 @@ func (r *MySQLRepository) ClaimNext(ctx context.Context, owner string, leaseDura return err } started := time.Now().UTC() - attempt := model.Attempt{StartedAt: &started} + attempt := model.Attempt{Type: "lease", StartedAt: &started} if generation.ProviderModelID != nil { attempt.ProviderModelID = *generation.ProviderModelID } @@ -141,14 +141,18 @@ func (r *MySQLRepository) ClaimNext(ctx context.Context, owner string, leaseDura return fmt.Errorf("encode expired queue attempt: %w", marshalErr) } attemptsExpression = `JSON_ARRAY_APPEND( - JSON_SET( - attempts, - CONCAT('$[', attempt_count - 1, ']'), - JSON_MERGE_PATCH( - JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), - CAST(? AS JSON) + CASE + WHEN JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].finished_at')) IS NULL + THEN JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, ']'), + JSON_MERGE_PATCH( + JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), + CAST(? AS JSON) + ) ) - ), + ELSE attempts + END, '$', CAST(? AS JSON) )` arguments = append(arguments, string(expiredAttempt), string(attemptJSON), generation.ID) @@ -181,6 +185,95 @@ func (r *MySQLRepository) ClaimNext(ctx context.Context, owner string, leaseDura return claim, err } +// BeginProviderAttempt appends one upstream-call event under the current +// generation lease. A stale worker receives owned=false and must not call an +// upstream provider. +func (r *MySQLRepository) BeginProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, routeMemberID, providerModelID uint64) (model.Attempt, bool, error) { + if generationID == 0 || leaseToken == "" || routeMemberID == 0 || providerModelID == 0 { + return model.Attempt{}, false, ErrInvalidAttempt + } + var attempt model.Attempt + owned := false + err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + var generation model.Generation + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&generation, generationID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil + } + return fmt.Errorf("lock generation for provider attempt: %w", err) + } + if generation.Status != model.StatusRunning || generation.LeaseToken == nil || *generation.LeaseToken != leaseToken { + return nil + } + started := time.Now().UTC() + attempt = model.Attempt{ + Type: "provider", RouteMemberID: routeMemberID, ProviderModelID: providerModelID, + ProviderOrdinal: generation.ProviderAttemptCount + 1, StartedAt: &started, + } + encoded, marshalErr := json.Marshal(attempt) + if marshalErr != nil { + return fmt.Errorf("encode provider attempt: %w", marshalErr) + } + result := tx.Exec(` + UPDATE generations + SET provider_model_id = ?, provider_attempt_count = provider_attempt_count + 1, + attempts = JSON_ARRAY_APPEND(attempts, '$', CAST(? AS JSON)), + attempt_count = attempt_count + 1 + WHERE id = ? AND status = 'running' AND lease_token = ?`, + providerModelID, string(encoded), generationID, leaseToken) + if result.Error != nil { + return fmt.Errorf("begin provider attempt: %w", result.Error) + } + owned = result.RowsAffected == 1 + return nil + }) + return attempt, owned, err +} + +// FinishProviderAttempt records the terminal observation for the current +// upstream call but leaves the generation running so retry logic can choose a +// different member. The lease token is part of the same CAS predicate. +func (r *MySQLRepository) FinishProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, attempt model.Attempt) (bool, error) { + if generationID == 0 || leaseToken == "" || attempt.RouteMemberID == 0 || attempt.ProviderModelID == 0 { + return false, ErrInvalidAttempt + } + attempt.Type = "provider" + attempt.ErrorCode = sanitizeCodeOrEmpty(attempt.ErrorCode) + attempt.ErrorMessage = SanitizeErrorMessage(attempt.ErrorMessage, 512) + if attempt.ErrorCode == "" { + attempt.ErrorMessage = "" + } + if attempt.LatencyMS < 0 { + attempt.LatencyMS = 0 + } + if attempt.FinishedAt == nil { + finished := time.Now().UTC() + attempt.FinishedAt = &finished + } + encoded, err := json.Marshal(attempt) + if err != nil { + return false, fmt.Errorf("encode finished provider attempt: %w", err) + } + result := r.db.WithContext(ctx).Exec(` + UPDATE generations + SET attempts = JSON_SET( + attempts, + CONCAT('$[', attempt_count - 1, ']'), + JSON_MERGE_PATCH( + JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, ']')), + CAST(? AS JSON) + ) + ) + WHERE id = ? AND status = 'running' AND lease_token = ? AND attempt_count > 0 + AND JSON_UNQUOTE(JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].type'))) = 'provider' + AND JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].finished_at')) IS NULL`, + string(encoded), generationID, leaseToken) + if result.Error != nil { + return false, fmt.Errorf("finish provider attempt: %w", result.Error) + } + return result.RowsAffected == 1, nil +} + func (r *MySQLRepository) AssignProvider(ctx context.Context, generationID uint64, leaseToken string, providerModelID uint64) (bool, error) { if generationID == 0 || leaseToken == "" || providerModelID == 0 { return false, ErrInvalidAttempt @@ -306,8 +399,7 @@ func (r *MySQLRepository) Fail(ctx context.Context, generationID uint64, leaseTo CAST(? AS JSON) ) ) - WHERE id = ? AND status = 'running' AND lease_token = ? - AND JSON_EXTRACT(attempts, CONCAT('$[', attempt_count - 1, '].provider_model_id')) IS NOT NULL`, + WHERE id = ? AND status = 'running' AND lease_token = ?`, code, message, string(attemptJSON), generationID, leaseToken) if result.Error != nil { return false, fmt.Errorf("fail queue generation: %w", result.Error) diff --git a/internal/core/queue/queue.go b/internal/core/queue/queue.go index c91966f..277b3c9 100644 --- a/internal/core/queue/queue.go +++ b/internal/core/queue/queue.go @@ -15,6 +15,8 @@ type Claim struct { type Repository interface { CreateIdempotent(ctx context.Context, generation *model.Generation, inputs []model.GenerationInput) (created bool, err error) ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*Claim, error) + BeginProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, routeMemberID, providerModelID uint64) (model.Attempt, bool, error) + FinishProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, attempt model.Attempt) (bool, error) AssignProvider(ctx context.Context, generationID uint64, leaseToken string, providerModelID uint64) (bool, error) Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempt model.Attempt) (bool, error) diff --git a/internal/core/queue/sanitize.go b/internal/core/queue/sanitize.go index 431501b..7aedb68 100644 --- a/internal/core/queue/sanitize.go +++ b/internal/core/queue/sanitize.go @@ -39,6 +39,13 @@ func sanitizeCode(code string) string { return code } +func sanitizeCodeOrEmpty(code string) string { + if strings.TrimSpace(code) == "" { + return "" + } + return sanitizeCode(code) +} + func truncateUTF8(value string, maxBytes int) string { if maxBytes <= 0 { return "" diff --git a/internal/core/router/random.go b/internal/core/router/random.go new file mode 100644 index 0000000..08dd571 --- /dev/null +++ b/internal/core/router/random.go @@ -0,0 +1,24 @@ +package router + +import ( + cryptorand "crypto/rand" + "math/big" +) + +// CryptoRandom is the production source for weighted selection. Returning the +// upper bound on a rare entropy failure makes SelectCandidates reject the +// value rather than silently biasing traffic to a route member. +type CryptoRandom struct{} + +func (CryptoRandom) Uint64n(limit uint64) uint64 { + if limit == 0 { + return 0 + } + value, err := cryptorand.Int(cryptorand.Reader, new(big.Int).SetUint64(limit)) + if err != nil { + return limit + } + return value.Uint64() +} + +var _ RandomSource = CryptoRandom{} diff --git a/internal/core/router/routing.go b/internal/core/router/routing.go index 6bbec15..99a04e1 100644 --- a/internal/core/router/routing.go +++ b/internal/core/router/routing.go @@ -30,6 +30,7 @@ var ( ErrRouteNotConfigured = Error{Code: CodeRouteNotConfigured} ErrRouteUnavailable = Error{Code: CodeRouteUnavailable} ErrFailoverExhausted = Error{Code: CodeFailoverExhausted} + ErrMemberUnavailable = errors.New("route member is unavailable") ErrInvalidSnapshot = errors.New("route snapshot is invalid") ErrInvalidRandom = errors.New("route random source is invalid") ) @@ -256,12 +257,14 @@ type SnapshotRepository interface { // #22. The owner/token lease lets implementations cap concurrent half-open // probes without weakening the circuit across worker processes. type RuntimeRepository interface { + MemberStates(context.Context, model.Capability, []MemberSnapshot) ([]MemberState, error) Reserve(context.Context, ReservationRequest) (Reservation, error) Record(context.Context, Reservation, CircuitObservation) (bool, error) } type ReservationRequest struct { RoutePoolMemberID uint64 + Capability model.Capability Owner string LeaseDuration time.Duration } @@ -273,7 +276,9 @@ type Reservation struct { } type CircuitObservation struct { - Retryable bool - Succeeded bool - ErrorCode string + Retryable bool + Succeeded bool + OpenImmediately bool + ErrorCode string + ErrorMessage string } diff --git a/internal/platform/mockprovider/handler.go b/internal/platform/mockprovider/handler.go index 3447765..439836e 100644 --- a/internal/platform/mockprovider/handler.go +++ b/internal/platform/mockprovider/handler.go @@ -7,6 +7,7 @@ import ( "image" "image/color" "image/png" + "mime/multipart" "net/http" "strings" "time" @@ -18,12 +19,18 @@ func (h Handler) ServeHTTP(response http.ResponseWriter, request *http.Request) switch request.URL.Path { case "/v1/chat/completions": h.chat(response, request) + case "/v1/images/generations": + h.image(response, request, false) case "/v1/images/edits": - h.image(response, request) + h.image(response, request, true) case "/v1/result.png": response.Header().Set("Content-Type", "image/png") response.Write(mockPNG()) default: + if strings.HasPrefix(request.URL.Path, "/v1/models/") && strings.HasSuffix(request.URL.Path, ":generateContent") { + h.gemini(response, request) + return + } http.NotFound(response, request) } } @@ -46,17 +53,30 @@ func (h Handler) chat(response http.ResponseWriter, request *http.Request) { _ = json.NewEncoder(response).Encode(map[string]any{"choices": []any{map[string]any{"message": map[string]string{"content": "mock: " + prompt}}}}) } -func (h Handler) image(response http.ResponseWriter, request *http.Request) { - if request.ParseMultipartForm(8<<20) != nil { - writeError(response, 400, "bad_request") - return +func (h Handler) image(response http.ResponseWriter, request *http.Request, requireInput bool) { + prompt := "" + var files []*multipart.FileHeader + if strings.HasPrefix(request.Header.Get("Content-Type"), "application/json") { + var body struct { + Prompt string `json:"prompt"` + } + if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil { + writeError(response, 400, "bad_request") + return + } + prompt = body.Prompt + } else { + if request.ParseMultipartForm(8<<20) != nil { + writeError(response, 400, "bad_request") + return + } + prompt = request.FormValue("prompt") + files = request.MultipartForm.File["image[]"] } - prompt := request.FormValue("prompt") if h.scenario(response, request, prompt) { return } - files := request.MultipartForm.File["image[]"] - if len(files) == 0 { + if requireInput && len(files) == 0 { writeError(response, 400, "bad_request") return } @@ -81,6 +101,43 @@ func (h Handler) image(response http.ResponseWriter, request *http.Request) { _ = json.NewEncoder(response).Encode(map[string]any{"data": []any{data}}) } +func (h Handler) gemini(response http.ResponseWriter, request *http.Request) { + var body struct { + Contents []struct { + Parts []struct { + Text string `json:"text"` + InlineData *struct { + MIMEType string `json:"mimeType"` + Data string `json:"data"` + } `json:"inlineData"` + } `json:"parts"` + } `json:"contents"` + } + if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil || len(body.Contents) != 1 { + writeError(response, 400, "bad_request") + return + } + prompt := "" + hasInlineData := false + for _, part := range body.Contents[0].Parts { + if part.Text != "" { + prompt = part.Text + } + if part.InlineData != nil && part.InlineData.MIMEType != "" && part.InlineData.Data != "" { + hasInlineData = true + } + } + if prompt == "" || h.scenario(response, request, prompt) { + return + } + part := map[string]any{"text": "mock: " + prompt} + if hasInlineData { + part = map[string]any{"inlineData": map[string]string{"mimeType": "image/png", "data": base64.StdEncoding.EncodeToString(mockPNG())}} + } + response.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(response).Encode(map[string]any{"candidates": []any{map[string]any{"content": map[string]any{"parts": []any{part}}}}}) +} + func (h Handler) scenario(response http.ResponseWriter, request *http.Request, prompt string) bool { switch { case strings.Contains(prompt, "mock:429"): diff --git a/internal/platform/mockprovider/handler_test.go b/internal/platform/mockprovider/handler_test.go index 236ebf6..85f3c17 100644 --- a/internal/platform/mockprovider/handler_test.go +++ b/internal/platform/mockprovider/handler_test.go @@ -41,3 +41,20 @@ func TestImageSuccess(t *testing.T) { t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) } } + +func TestImagesGenerationAndGeminiEndpoints(t *testing.T) { + imageRequest := httptest.NewRequest(http.MethodPost, "http://mock/v1/images/generations", bytes.NewBufferString(`{"model":"mock","prompt":"draw","n":1}`)) + imageRequest.Header.Set("Content-Type", "application/json") + imageResponse := httptest.NewRecorder() + Handler{}.ServeHTTP(imageResponse, imageRequest) + if imageResponse.Code != http.StatusOK || !bytes.Contains(imageResponse.Body.Bytes(), []byte("b64_json")) { + t.Fatalf("images status=%d body=%s", imageResponse.Code, imageResponse.Body.String()) + } + + geminiRequest := httptest.NewRequest(http.MethodPost, "http://mock/v1/models/mock:generateContent", bytes.NewBufferString(`{"contents":[{"parts":[{"text":"hello"}]}]}`)) + geminiResponse := httptest.NewRecorder() + Handler{}.ServeHTTP(geminiResponse, geminiRequest) + if geminiResponse.Code != http.StatusOK || !bytes.Contains(geminiResponse.Body.Bytes(), []byte("candidates")) { + t.Fatalf("gemini status=%d body=%s", geminiResponse.Code, geminiResponse.Body.String()) + } +} diff --git a/portal/handler/mysql_integration_test.go b/portal/handler/mysql_integration_test.go index c2934eb..65048d5 100644 --- a/portal/handler/mysql_integration_test.go +++ b/portal/handler/mysql_integration_test.go @@ -107,8 +107,16 @@ func TestPortalAuthenticationSubmissionAndAuthorization(t *testing.T) { if err := db.Create(&users).Error; err != nil { t.Fatal(err) } + templates := []model.PromptTemplate{ + {TemplateKey: "portal-text-" + suffix, Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Portal Text", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true}, + {TemplateKey: "portal-image-" + suffix, Kind: model.KindImage, APIType: model.APIImagesEdits, Capability: model.CapabilityImageEdit, Name: "Portal Image", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true}, + } + if err := db.Create(&templates).Error; err != nil { + t.Fatal(err) + } defer func() { db.Where("user_id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.Generation{}) + db.Where("id IN ?", []uint64{templates[0].ID, templates[1].ID}).Delete(&model.PromptTemplate{}) db.Where("id IN ?", []uint64{users[0].ID, users[1].ID}).Delete(&model.User{}) }() queueRepository, _ := queue.NewMySQLRepository(db) diff --git a/portal/main.go b/portal/main.go index bec32ea..210ec24 100644 --- a/portal/main.go +++ b/portal/main.go @@ -91,6 +91,10 @@ func run() error { if err != nil { return err } + runtime, err := worker.NewGORMRuntime(db) + if err != nil { + return err + } workerStorage, err := worker.NewLocalStorage(storage) if err != nil { return err @@ -99,7 +103,7 @@ func run() error { Owner: fmt.Sprintf("portal-%d", os.Getpid()), LeaseDuration: cfg.WorkerLeaseDuration, PollInterval: cfg.WorkerPollInterval, - }, queueController, queueController, catalog, providerFactory, workerStorage) + }, queueController, queueController, catalog, providerFactory, runtime, workerStorage) if err != nil { return err } diff --git a/portal/worker/catalog.go b/portal/worker/catalog.go index 2434165..6be3e55 100644 --- a/portal/worker/catalog.go +++ b/portal/worker/catalog.go @@ -9,6 +9,7 @@ import ( corecrypto "git.ilapage.cn/OPC/chorus/internal/core/crypto" "git.ilapage.cn/OPC/chorus/internal/core/model" "git.ilapage.cn/OPC/chorus/internal/core/provider" + "git.ilapage.cn/OPC/chorus/internal/core/router" "gorm.io/gorm" ) @@ -24,28 +25,48 @@ func NewGORMCatalog(db *gorm.DB, cipher corecrypto.KeyCipher) (*GORMCatalog, err return &GORMCatalog{db: db, cipher: cipher}, nil } -func (c *GORMCatalog) Select(ctx context.Context, kind model.GenerationKind) (Selection, error) { - type row struct { - ProviderModelID uint64 - BaseURL, AuthType, ModelID string - APIKeyEnc json.RawMessage - APIType model.APIType - ExtraBody json.RawMessage - TimeoutMS uint32 - } - var rows []row - err := c.db.WithContext(ctx).Table("provider_models pm").Select("pm.id provider_model_id, p.base_url, p.auth_type, p.api_key_enc, pm.model_id, pm.api_type, pm.extra_body, pm.timeout_ms").Joins("JOIN providers p ON p.id = pm.provider_id").Where("pm.kind = ? AND pm.enabled = TRUE AND p.enabled = TRUE", kind).Order("pm.id").Limit(2).Scan(&rows).Error - if err != nil { - return Selection{}, fmt.Errorf("select enabled provider model: %w", err) - } - if len(rows) != 1 { +// Resolve reads the current executable configuration only after the route +// member has been selected. The route snapshot never contains these fields. +func (c *GORMCatalog) Resolve(ctx context.Context, member router.MemberSnapshot, capability model.Capability) (Selection, error) { + if member.RoutePoolMemberID == 0 || member.ProviderModelID == 0 || !capability.Valid() { + return Selection{}, ErrNoProvider + } + type row struct { + ProviderModelID uint64 + BaseURL string + AuthType string + APIKeyEnc json.RawMessage + ModelID string + APIType model.APIType + ExtraBody json.RawMessage + TimeoutMS uint32 + } + var selected row + result := c.db.WithContext(ctx).Raw(` + SELECT pm.id AS provider_model_id, p.base_url, p.auth_type, credential.api_key_enc, + pm.model_id, pm.api_type, pm.extra_body, pm.timeout_ms + FROM route_pool_members member + JOIN provider_models pm ON pm.id = member.provider_model_id + JOIN providers p ON p.id = pm.provider_id + LEFT JOIN provider_credentials credential + ON credential.id = p.active_credential_id AND credential.status = 'active' + WHERE member.id = ? AND member.provider_model_id = ? + AND member.enabled = TRUE AND pm.enabled = TRUE AND p.enabled = TRUE + AND EXISTS(SELECT 1 FROM provider_model_capabilities capability_row + WHERE capability_row.provider_model_id = pm.id AND capability_row.capability = ?) + LIMIT 1`, member.RoutePoolMemberID, member.ProviderModelID, capability).Scan(&selected) + if result.Error != nil { + return Selection{}, fmt.Errorf("load routed provider model: %w", result.Error) + } + if result.RowsAffected != 1 { + return Selection{}, ErrNoProvider + } + authType := provider.AuthType(selected.AuthType) + if !authType.Valid() { return Selection{}, ErrNoProvider } - selected := rows[0] var apiKey string - switch selected.AuthType { - case "none": - case "bearer": + if authType != provider.AuthNone { var envelope corecrypto.Envelope if json.Unmarshal(selected.APIKeyEnc, &envelope) != nil { return Selection{}, fmt.Errorf("decode provider credential") @@ -55,13 +76,15 @@ func (c *GORMCatalog) Select(ctx context.Context, kind model.GenerationKind) (Se return Selection{}, fmt.Errorf("decrypt provider credential: %w", err) } apiKey = string(plaintext) - for i := range plaintext { - plaintext[i] = 0 + for index := range plaintext { + plaintext[index] = 0 } - default: - return Selection{}, fmt.Errorf("unsupported provider authentication") } - return Selection{ProviderModelID: selected.ProviderModelID, BaseURL: selected.BaseURL, APIKey: apiKey, ModelID: selected.ModelID, APIType: selected.APIType, ExtraBody: selected.ExtraBody, Timeout: time.Duration(selected.TimeoutMS) * time.Millisecond}, nil + return Selection{ + ProviderModelID: selected.ProviderModelID, BaseURL: selected.BaseURL, AuthType: authType, + APIKey: apiKey, ModelID: selected.ModelID, APIType: selected.APIType, + ExtraBody: selected.ExtraBody, Timeout: time.Duration(selected.TimeoutMS) * time.Millisecond, + }, nil } type OpenAIFactory struct { @@ -75,6 +98,10 @@ func NewOpenAIFactory(httpClient provider.HTTPClient, maxResponseBytes int64) (* } return &OpenAIFactory{http: httpClient, maxResponseBytes: maxResponseBytes}, nil } + func (f *OpenAIFactory) New(selection Selection) (provider.Client, error) { - return provider.NewOpenAI(f.http, provider.OpenAIConfig{BaseURL: selection.BaseURL, APIKey: selection.APIKey, ExtraBody: selection.ExtraBody, MaxResponseBytes: f.maxResponseBytes}) + return provider.NewOpenAI(f.http, provider.OpenAIConfig{ + BaseURL: selection.BaseURL, AuthType: selection.AuthType, APIKey: selection.APIKey, + ExtraBody: selection.ExtraBody, MaxResponseBytes: f.maxResponseBytes, + }) } diff --git a/portal/worker/mysql_integration_test.go b/portal/worker/mysql_integration_test.go index b73bf7d..a24866f 100644 --- a/portal/worker/mysql_integration_test.go +++ b/portal/worker/mysql_integration_test.go @@ -4,23 +4,24 @@ import ( "bytes" "context" "encoding/json" + "errors" "fmt" "net" "net/http/httptest" "os" - "strconv" "testing" "time" "git.ilapage.cn/OPC/chorus/internal/core/model" "git.ilapage.cn/OPC/chorus/internal/core/queue" - corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage" + "git.ilapage.cn/OPC/chorus/internal/core/router" platformcrypto "git.ilapage.cn/OPC/chorus/internal/platform/crypto" safehttp "git.ilapage.cn/OPC/chorus/internal/platform/http" "git.ilapage.cn/OPC/chorus/internal/platform/mockprovider" platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage" "gorm.io/driver/mysql" "gorm.io/gorm" + "gorm.io/gorm/logger" ) type fixedResolver struct{} @@ -29,29 +30,38 @@ func (fixedResolver) LookupIPAddr(context.Context, string) ([]net.IPAddr, error) return []net.IPAddr{{IP: net.ParseIP("93.184.216.34")}}, nil } -func TestMySQLWorkerWithMockUpstream(t *testing.T) { +func TestMySQLWorkerUsesSnapshotAndMockUpstream(t *testing.T) { dsn := os.Getenv("CHORUS_TEST_DSN") if dsn == "" { t.Skip("CHORUS_TEST_DSN is not set") } - db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{}) + db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) if err != nil { t.Fatal(err) } - sqlDB, _ := db.DB() - t.Cleanup(func() { sqlDB.Close() }) + sqlDB, err := db.DB() + if err != nil { + t.Fatal(err) + } + defer sqlDB.Close() + tx := db.Begin() + if tx.Error != nil { + t.Fatal(tx.Error) + } + defer tx.Rollback() + server := httptest.NewServer(mockprovider.Handler{}) defer server.Close() - address := server.Listener.Addr().String() - httpClient, err := safehttp.New(safehttp.Config{Resolver: fixedResolver{}, DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) { - var dialer net.Dialer - return dialer.DialContext(ctx, network, address) - }, Timeout: 2 * time.Second, MaxRedirects: 1, AllowedPorts: []uint16{80}}) + httpClient, err := safehttp.New(safehttp.Config{ + Resolver: fixedResolver{}, Timeout: 2 * time.Second, MaxRedirects: 1, AllowedPorts: []uint16{80}, + DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, network, server.Listener.Addr().String()) + }, + }) if err != nil { t.Fatal(err) } - providerHTTP := safehttp.NewProviderClient(httpClient) - factory, err := NewOpenAIFactory(providerHTTP, 2<<20) + factory, err := NewOpenAIFactory(safehttp.NewProviderClient(httpClient), 2<<20) if err != nil { t.Fatal(err) } @@ -59,39 +69,15 @@ func TestMySQLWorkerWithMockUpstream(t *testing.T) { if err != nil { t.Fatal(err) } - catalog, err := NewGORMCatalog(db, keyRing) + catalog, err := NewGORMCatalog(tx, keyRing) if err != nil { t.Fatal(err) } - var providerID uint64 - var oldURL, oldAuth string - var oldKey []byte - if err := db.Raw("SELECT id, base_url, auth_type, api_key_enc FROM providers WHERE enabled=TRUE ORDER BY id LIMIT 1").Row().Scan(&providerID, &oldURL, &oldAuth, &oldKey); err != nil { - t.Fatal(err) - } - envelope, err := keyRing.Encrypt(context.Background(), []byte("integration-provider-key")) + runtime, err := NewGORMRuntime(tx) if err != nil { t.Fatal(err) } - encryptedKey, err := json.Marshal(envelope) - if err != nil { - t.Fatal(err) - } - if err := db.Exec("UPDATE providers SET base_url=?, auth_type='bearer', api_key_enc=? WHERE id=?", "http://provider.test/v1", string(encryptedKey), providerID).Error; err != nil { - t.Fatal(err) - } - defer func() { - if oldKey == nil { - db.Exec("UPDATE providers SET base_url=?, auth_type=?, api_key_enc=NULL WHERE id=?", oldURL, oldAuth, providerID) - } else { - db.Exec("UPDATE providers SET base_url=?, auth_type=?, api_key_enc=? WHERE id=?", oldURL, oldAuth, string(oldKey), providerID) - } - }() - selected, err := catalog.Select(context.Background(), model.KindText) - if err != nil || selected.APIKey != "integration-provider-key" { - t.Fatalf("encrypted provider key was not selected correctly: %v", err) - } - queueRepository, err := queue.NewMySQLRepository(db) + queueRepository, err := queue.NewMySQLRepository(tx) if err != nil { t.Fatal(err) } @@ -99,59 +85,174 @@ func TestMySQLWorkerWithMockUpstream(t *testing.T) { if err != nil { t.Fatal(err) } - store, _ := NewLocalStorage(local) - var userID uint64 - if err := db.Raw("SELECT id FROM users WHERE status='active' ORDER BY id LIMIT 1").Scan(&userID).Error; err != nil || userID == 0 { - t.Fatalf("seed user unavailable: %v", err) + store, err := NewLocalStorage(local) + if err != nil { + t.Fatal(err) } - for _, kind := range []model.GenerationKind{model.KindText, model.KindImage} { - t.Run(string(kind), func(t *testing.T) { - prompt := "integration" - if kind == model.KindImage { - prompt = "mock:url" - } - // Keep this fixture ahead of unrelated pending rows in a shared integration database. - generation := model.Generation{UserID: userID, Kind: kind, IdempotencyKey: "worker-it-" + string(kind) + "-" + strconv.FormatInt(time.Now().UnixNano(), 10), UserPrompt: prompt, RenderedPrompt: prompt, CreatedAt: time.Unix(1, 0)} - created, createErr := queueRepository.CreateIdempotent(context.Background(), &generation, nil) - if createErr != nil || !created { - t.Fatalf("create=%v err=%v", created, createErr) - } - defer db.Exec("DELETE FROM generations WHERE id=?", generation.ID) - if kind == model.KindImage { - inputKey := fmt.Sprintf("integration/%d/input", time.Now().UnixNano()) - data := testPNG() - object, putErr := local.Put(context.Background(), corestorage.PutRequest{Key: inputKey, OwnerID: userID, GenerationID: generation.ID, ContentType: "image/png", Source: bytes.NewReader(data)}) - if putErr != nil { - t.Fatal(putErr) - } - input := model.GenerationInput{GenerationID: generation.ID, Position: 0, Role: model.RolePrimary, OriginalName: "input.png", MIMEType: "image/png", StorageKey: object.Key, SizeBytes: uint64(object.Size)} - if err := db.Create(&input).Error; err != nil { - t.Fatal(err) - } - } - controller := queue.NewController(queueRepository) - w, err := New(Config{Owner: "integration", LeaseDuration: 5 * time.Second, PollInterval: time.Millisecond}, queueRepository, controller, catalog, factory, store) - if err != nil { - t.Fatal(err) - } - worked, err := w.ProcessOne(context.Background()) - if err != nil || !worked { - t.Fatalf("worked=%v err=%v", worked, err) - } - var status string - var attemptCount int - var attempts string - if err := db.Raw("SELECT status, attempt_count, CAST(attempts AS CHAR) FROM generations WHERE id=?", generation.ID).Row().Scan(&status, &attemptCount, &attempts); err != nil { - t.Fatal(err) - } - if status != "succeeded" || attemptCount != 1 || !bytes.Contains([]byte(attempts), []byte("provider_model_id")) || !bytes.Contains([]byte(attempts), []byte("latency_ms")) { - t.Fatalf("status=%s count=%d attempts=%s", status, attemptCount, attempts) - } - var outputCount int64 - db.Model(&model.GenerationOutput{}).Where("generation_id=?", generation.ID).Count(&outputCount) - if outputCount != 1 { - t.Fatalf("outputs=%d", outputCount) - } - }) + + suffix := time.Now().UnixNano() + user := model.User{Email: fmt.Sprintf("worker-%d@chorus.invalid", suffix), PasswordHash: "synthetic", DisplayName: "Worker", Status: "active"} + providerRow := model.Provider{Slug: fmt.Sprintf("worker-%d", suffix), Name: "Worker Mock", BaseURL: "http://provider.test/v1", AuthType: "none", Enabled: true} + modelRow := model.ProviderModel{ProviderID: 0, Name: "Worker Chat", ModelID: "mock-chat", APIType: model.APIChat, Kind: model.KindText, ExtraBody: json.RawMessage("{}"), TimeoutMS: 1000, Weight: 1, Enabled: true} + template := model.PromptTemplate{TemplateKey: fmt.Sprintf("worker-%d", suffix), Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Worker", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true} + if err := tx.Create(&user).Error; err != nil { + t.Fatal(err) + } + if err := tx.Create(&providerRow).Error; err != nil { + t.Fatal(err) + } + modelRow.ProviderID = providerRow.ID + if err := tx.Create(&modelRow).Error; err != nil { + t.Fatal(err) + } + if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id, capability) VALUES (?, ?)", modelRow.ID, model.CapabilityText).Error; err != nil { + t.Fatal(err) + } + if err := tx.Create(&template).Error; err != nil { + t.Fatal(err) + } + if err := tx.Exec("INSERT INTO route_pools (slug, name, capability, prompt_template_id, max_failover, version, enabled) VALUES (?, ?, ?, ?, 0, 1, TRUE)", fmt.Sprintf("worker-%d", suffix), "Worker", model.CapabilityText, template.ID).Error; err != nil { + t.Fatal(err) + } + var poolID uint64 + if err := tx.Raw("SELECT id FROM route_pools WHERE slug = ?", fmt.Sprintf("worker-%d", suffix)).Scan(&poolID).Error; err != nil || poolID == 0 { + t.Fatalf("route pool id=%d error=%v", poolID, err) + } + if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id, provider_model_id, weight, failure_threshold, open_seconds, half_open_max, enabled) VALUES (?, ?, 1, 2, 60, 1, TRUE)", poolID, modelRow.ID).Error; err != nil { + t.Fatal(err) + } + var memberID uint64 + if err := tx.Raw("SELECT id FROM route_pool_members WHERE route_pool_id = ?", poolID).Scan(&memberID).Error; err != nil || memberID == 0 { + t.Fatalf("route member id=%d error=%v", memberID, err) + } + + snapshot := router.RouteSnapshot{Capability: model.CapabilityText, RoutePoolID: poolID, RoutePoolVersion: 1, PromptTemplateID: template.ID, PromptTemplateKey: template.TemplateKey, PromptTemplateVersion: 1, Members: []router.MemberSnapshot{{RoutePoolMemberID: memberID, ProviderModelID: modelRow.ID, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}}} + generation := model.Generation{UserID: user.ID, Kind: model.KindText, IdempotencyKey: fmt.Sprintf("worker-%d", suffix), UserPrompt: "integration", RenderedPrompt: "integration", CreatedAt: time.Unix(1, 0)} + if err := router.ApplySnapshot(&generation, snapshot); err != nil { + t.Fatal(err) + } + created, err := queueRepository.CreateIdempotent(context.Background(), &generation, nil) + if err != nil || !created { + t.Fatalf("create=%v error=%v", created, err) + } + controller := queue.NewController(queueRepository) + w, err := New(Config{Owner: "integration", LeaseDuration: 5 * time.Second, PollInterval: time.Millisecond, Random: firstRandom{}}, queueRepository, controller, catalog, factory, runtime, store) + if err != nil { + t.Fatal(err) + } + worked, err := w.ProcessOne(context.Background()) + if err != nil || !worked { + t.Fatalf("worked=%v error=%v", worked, err) + } + var stored model.Generation + if err := tx.First(&stored, generation.ID).Error; err != nil { + t.Fatal(err) + } + if stored.Status != model.StatusSucceeded || stored.ProviderAttemptCount != 1 || !bytes.Contains(stored.Attempts, []byte(`"route_member_id"`)) || !bytes.Contains(stored.Attempts, []byte(`"finished_at"`)) { + t.Fatalf("stored generation=%#v attempts=%s", stored, stored.Attempts) + } + var outputCount int64 + if err := tx.Model(&model.GenerationOutput{}).Where("generation_id = ?", generation.ID).Count(&outputCount).Error; err != nil || outputCount != 1 { + t.Fatalf("outputs=%d error=%v", outputCount, err) + } +} + +func TestMySQLRuntimeOpensAndLimitsHalfOpenProbe(t *testing.T) { + dsn := os.Getenv("CHORUS_TEST_DSN") + if dsn == "" { + t.Skip("CHORUS_TEST_DSN is not set") + } + db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}) + if err != nil { + t.Fatal(err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatal(err) + } + defer sqlDB.Close() + tx := db.Begin() + if tx.Error != nil { + t.Fatal(tx.Error) + } + defer tx.Rollback() + + suffix := time.Now().UnixNano() + providerRow := model.Provider{Slug: fmt.Sprintf("runtime-%d", suffix), Name: "Runtime", BaseURL: "https://provider.invalid/v1", AuthType: "none", Enabled: true} + if err := tx.Create(&providerRow).Error; err != nil { + t.Fatal(err) + } + modelRow := model.ProviderModel{ProviderID: providerRow.ID, Name: "Runtime", ModelID: "runtime", APIType: model.APIChat, Kind: model.KindText, ExtraBody: json.RawMessage("{}"), TimeoutMS: 1000, Weight: 1, Enabled: true} + if err := tx.Create(&modelRow).Error; err != nil { + t.Fatal(err) + } + if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id, capability) VALUES (?, ?)", modelRow.ID, model.CapabilityText).Error; err != nil { + t.Fatal(err) + } + template := model.PromptTemplate{TemplateKey: fmt.Sprintf("runtime-%d", suffix), Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Runtime", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true} + if err := tx.Create(&template).Error; err != nil { + t.Fatal(err) + } + poolSlug := fmt.Sprintf("runtime-%d", suffix) + if err := tx.Exec("INSERT INTO route_pools (slug, name, capability, prompt_template_id, max_failover, version, enabled) VALUES (?, 'Runtime', ?, ?, 0, 1, TRUE)", poolSlug, model.CapabilityText, template.ID).Error; err != nil { + t.Fatal(err) + } + var poolID uint64 + if err := tx.Raw("SELECT id FROM route_pools WHERE slug = ?", poolSlug).Scan(&poolID).Error; err != nil || poolID == 0 { + t.Fatalf("pool id=%d error=%v", poolID, err) + } + if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id, provider_model_id, weight, failure_threshold, open_seconds, half_open_max, enabled) VALUES (?, ?, 1, 2, 60, 1, TRUE)", poolID, modelRow.ID).Error; err != nil { + t.Fatal(err) + } + var memberID uint64 + if err := tx.Raw("SELECT id FROM route_pool_members WHERE route_pool_id = ?", poolID).Scan(&memberID).Error; err != nil || memberID == 0 { + t.Fatalf("member id=%d error=%v", memberID, err) + } + runtime, err := NewGORMRuntime(tx) + if err != nil { + t.Fatal(err) + } + members := []router.MemberSnapshot{{RoutePoolMemberID: memberID, ProviderModelID: modelRow.ID, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}} + state := func() router.MemberState { + states, stateErr := runtime.MemberStates(context.Background(), model.CapabilityText, members) + if stateErr != nil || len(states) != 1 { + t.Fatalf("states=%#v error=%v", states, stateErr) + } + return states[0] + } + if current := state(); current.CircuitState != router.CircuitClosed || !current.Eligible() { + t.Fatalf("initial state=%#v", current) + } + for attempt := 0; attempt < 2; attempt++ { + reservation, reserveErr := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "runtime-test", LeaseDuration: time.Second}) + if reserveErr != nil || reservation.State != router.CircuitClosed { + t.Fatalf("closed reservation=%#v error=%v", reservation, reserveErr) + } + owned, recordErr := runtime.Record(context.Background(), reservation, router.CircuitObservation{Retryable: true, ErrorCode: "upstream_timeout", ErrorMessage: "upstream request failed"}) + if recordErr != nil || !owned { + t.Fatalf("record retryable owned=%v error=%v", owned, recordErr) + } + } + if current := state(); current.CircuitState != router.CircuitOpen || current.Eligible() { + t.Fatalf("open state=%#v", current) + } + if err := tx.Exec("UPDATE route_member_runtime SET open_until = ? WHERE route_pool_member_id = ?", time.Now().Add(-time.Second), memberID).Error; err != nil { + t.Fatal(err) + } + if current := state(); current.CircuitState != router.CircuitHalfOpen || current.HalfOpenInFlight != 0 || !current.Eligible() { + t.Fatalf("ready half-open state=%#v", current) + } + probe, err := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "runtime-test", LeaseDuration: time.Second}) + if err != nil || probe.State != router.CircuitHalfOpen || probe.Token == "" { + t.Fatalf("probe=%#v error=%v", probe, err) + } + if _, err := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "second-worker", LeaseDuration: time.Second}); !errors.Is(err, router.ErrMemberUnavailable) { + t.Fatalf("second probe error=%v", err) + } + if owned, err := runtime.Record(context.Background(), probe, router.CircuitObservation{Succeeded: true}); err != nil || !owned { + t.Fatalf("successful probe owned=%v error=%v", owned, err) + } + if current := state(); current.CircuitState != router.CircuitClosed || current.HalfOpenInFlight != 0 || !current.Eligible() { + t.Fatalf("closed after probe state=%#v", current) } } diff --git a/portal/worker/runtime.go b/portal/worker/runtime.go new file mode 100644 index 0000000..0825697 --- /dev/null +++ b/portal/worker/runtime.go @@ -0,0 +1,289 @@ +package worker + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "strings" + "time" + + "git.ilapage.cn/OPC/chorus/internal/core/model" + "git.ilapage.cn/OPC/chorus/internal/core/queue" + "git.ilapage.cn/OPC/chorus/internal/core/router" + "gorm.io/gorm" +) + +// GORMRuntime is the cross-worker circuit breaker. It intentionally queries +// only runtime and enabled/capability flags; credentials and provider URLs are +// never part of this read path. +type GORMRuntime struct{ db *gorm.DB } + +func NewGORMRuntime(db *gorm.DB) (*GORMRuntime, error) { + if db == nil { + return nil, ErrInvalidConfig + } + return &GORMRuntime{db: db}, nil +} + +func (r *GORMRuntime) MemberStates(ctx context.Context, capability model.Capability, members []router.MemberSnapshot) ([]router.MemberState, error) { + if !capability.Valid() || len(members) == 0 { + return nil, router.ErrInvalidSnapshot + } + memberIDs := make([]uint64, 0, len(members)) + states := make(map[uint64]router.MemberState, len(members)) + for _, member := range members { + if member.RoutePoolMemberID == 0 { + return nil, router.ErrInvalidSnapshot + } + memberIDs = append(memberIDs, member.RoutePoolMemberID) + states[member.RoutePoolMemberID] = router.MemberState{RoutePoolMemberID: member.RoutePoolMemberID} + } + type row struct { + RoutePoolMemberID uint64 + MemberEnabled bool + ProviderEnabled bool + ModelEnabled bool + SupportsCapability bool + CircuitState string + OpenUntil *time.Time + ProbeUntil *time.Time + HalfOpenMax uint16 + } + var rows []row + err := r.db.WithContext(ctx).Raw(` + SELECT rpm.id AS route_pool_member_id, + rpm.enabled AS member_enabled, + p.enabled AS provider_enabled, + pm.enabled AS model_enabled, + EXISTS(SELECT 1 FROM provider_model_capabilities c + WHERE c.provider_model_id = pm.id AND c.capability = ?) AS supports_capability, + COALESCE(runtime.state, 'closed') AS circuit_state, + runtime.open_until, runtime.probe_until, rpm.half_open_max + FROM route_pool_members rpm + JOIN provider_models pm ON pm.id = rpm.provider_model_id + JOIN providers p ON p.id = pm.provider_id + LEFT JOIN route_member_runtime runtime ON runtime.route_pool_member_id = rpm.id + WHERE rpm.id IN ?`, capability, memberIDs).Scan(&rows).Error + if err != nil { + return nil, fmt.Errorf("load route member states: %w", err) + } + now := time.Now().UTC() + for _, row := range rows { + state := router.CircuitState(row.CircuitState) + halfOpenInFlight := uint16(0) + switch state { + case router.CircuitClosed: + case router.CircuitOpen: + if row.OpenUntil != nil && !row.OpenUntil.After(now) { + state = router.CircuitHalfOpen + } + case router.CircuitHalfOpen: + if row.ProbeUntil != nil && row.ProbeUntil.After(now) { + halfOpenInFlight = 1 + } + default: + state = router.CircuitOpen + } + states[row.RoutePoolMemberID] = router.MemberState{ + RoutePoolMemberID: row.RoutePoolMemberID, Enabled: row.MemberEnabled, + ProviderEnabled: row.ProviderEnabled, ModelEnabled: row.ModelEnabled, + SupportsCapability: row.SupportsCapability, CircuitState: state, + HalfOpenInFlight: halfOpenInFlight, HalfOpenMax: row.HalfOpenMax, + } + } + result := make([]router.MemberState, 0, len(memberIDs)) + for _, memberID := range memberIDs { + result = append(result, states[memberID]) + } + return result, nil +} + +func (r *GORMRuntime) Reserve(ctx context.Context, request router.ReservationRequest) (router.Reservation, error) { + if request.RoutePoolMemberID == 0 || !request.Capability.Valid() || strings.TrimSpace(request.Owner) == "" || len(request.Owner) > 128 || request.LeaseDuration <= 0 { + return router.Reservation{}, router.ErrMemberUnavailable + } + var reservation router.Reservation + err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Exec(`INSERT IGNORE INTO route_member_runtime (route_pool_member_id, state, consecutive_failures) VALUES (?, 'closed', 0)`, request.RoutePoolMemberID).Error; err != nil { + return fmt.Errorf("create route member runtime: %w", err) + } + type row struct { + MemberEnabled bool + ProviderEnabled bool + ModelEnabled bool + SupportsCapability bool + State string + OpenUntil *time.Time + ProbeToken *string + ProbeUntil *time.Time + } + var value row + result := tx.Raw(` + SELECT rpm.enabled AS member_enabled, p.enabled AS provider_enabled, pm.enabled AS model_enabled, + EXISTS(SELECT 1 FROM provider_model_capabilities c + WHERE c.provider_model_id = pm.id AND c.capability = ?) AS supports_capability, + runtime.state, runtime.open_until, runtime.probe_token, runtime.probe_until + FROM route_pool_members rpm + JOIN provider_models pm ON pm.id = rpm.provider_model_id + JOIN providers p ON p.id = pm.provider_id + JOIN route_member_runtime runtime ON runtime.route_pool_member_id = rpm.id + WHERE rpm.id = ? FOR UPDATE`, request.Capability, request.RoutePoolMemberID).Scan(&value) + if result.Error != nil { + return fmt.Errorf("lock route member runtime: %w", result.Error) + } + if result.RowsAffected != 1 || !value.MemberEnabled || !value.ProviderEnabled || !value.ModelEnabled || !value.SupportsCapability { + return router.ErrMemberUnavailable + } + now := time.Now().UTC() + state := router.CircuitState(value.State) + switch state { + case router.CircuitClosed: + reservation = router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, State: router.CircuitClosed} + return nil + case router.CircuitOpen: + if value.OpenUntil == nil || value.OpenUntil.After(now) { + return router.ErrMemberUnavailable + } + case router.CircuitHalfOpen: + if value.ProbeUntil != nil && value.ProbeUntil.After(now) { + return router.ErrMemberUnavailable + } + default: + return router.ErrMemberUnavailable + } + token, tokenErr := newRuntimeToken() + if tokenErr != nil { + return tokenErr + } + until := now.Add(request.LeaseDuration) + update := tx.Model(&runtimeRow{}).Where("route_pool_member_id = ?", request.RoutePoolMemberID).Updates(map[string]any{ + "state": router.CircuitHalfOpen, "open_until": nil, "probe_owner": request.Owner, + "probe_token": token, "probe_until": until, + }) + if update.Error != nil { + return fmt.Errorf("reserve half-open route member: %w", update.Error) + } + if update.RowsAffected != 1 { + return router.ErrMemberUnavailable + } + reservation = router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, Token: token, State: router.CircuitHalfOpen} + return nil + }) + return reservation, err +} + +func (r *GORMRuntime) Record(ctx context.Context, reservation router.Reservation, observation router.CircuitObservation) (bool, error) { + if reservation.RoutePoolMemberID == 0 { + return false, router.ErrMemberUnavailable + } + owned := false + err := r.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { + type row struct { + State string + ConsecutiveFailures uint32 + ProbeToken *string + ProbeUntil *time.Time + FailureThreshold uint32 + OpenSeconds uint32 + } + var value row + result := tx.Raw(` + SELECT runtime.state, runtime.consecutive_failures, runtime.probe_token, runtime.probe_until, + rpm.failure_threshold, rpm.open_seconds + FROM route_member_runtime runtime + JOIN route_pool_members rpm ON rpm.id = runtime.route_pool_member_id + WHERE runtime.route_pool_member_id = ? FOR UPDATE`, reservation.RoutePoolMemberID).Scan(&value) + if result.Error != nil { + return fmt.Errorf("lock route observation runtime: %w", result.Error) + } + if result.RowsAffected != 1 { + return nil + } + if reservation.State == router.CircuitHalfOpen { + if value.State != string(router.CircuitHalfOpen) || value.ProbeToken == nil || *value.ProbeToken != reservation.Token || value.ProbeUntil == nil || !value.ProbeUntil.After(time.Now().UTC()) { + return nil + } + } + updates := map[string]any{ + "last_observed_at": time.Now().UTC(), + "last_error_code": nullableRuntimeCode(observation.ErrorCode), + "last_error_message": nullableRuntimeMessage(observation.ErrorMessage), + } + switch { + case observation.Succeeded: + updates["state"] = router.CircuitClosed + updates["consecutive_failures"] = 0 + updates["open_until"] = nil + case observation.OpenImmediately: + updates["state"] = router.CircuitOpen + updates["consecutive_failures"] = value.FailureThreshold + updates["open_until"] = time.Now().UTC().Add(time.Duration(value.OpenSeconds) * time.Second) + case observation.Retryable: + next := value.ConsecutiveFailures + 1 + updates["consecutive_failures"] = next + if next >= value.FailureThreshold { + updates["state"] = router.CircuitOpen + updates["open_until"] = time.Now().UTC().Add(time.Duration(value.OpenSeconds) * time.Second) + } else { + updates["state"] = router.CircuitClosed + updates["open_until"] = nil + } + default: + updates["state"] = router.CircuitClosed + updates["consecutive_failures"] = 0 + updates["open_until"] = nil + } + updates["probe_owner"] = nil + updates["probe_token"] = nil + updates["probe_until"] = nil + update := tx.Model(&runtimeRow{}).Where("route_pool_member_id = ?", reservation.RoutePoolMemberID).Updates(updates) + if update.Error != nil { + return fmt.Errorf("record route observation: %w", update.Error) + } + owned = update.RowsAffected == 1 + return nil + }) + return owned, err +} + +type runtimeRow struct { + RoutePoolMemberID uint64 `gorm:"column:route_pool_member_id;primaryKey"` +} + +func (runtimeRow) TableName() string { return "route_member_runtime" } + +func nullableRuntimeCode(value string) any { + value = strings.ToLower(strings.TrimSpace(value)) + if value == "" { + return nil + } + if len(value) > 64 { + return "upstream_error" + } + for _, character := range value { + if character != '_' && (character < 'a' || character > 'z') && (character < '0' || character > '9') { + return "upstream_error" + } + } + return value +} + +func nullableRuntimeMessage(value string) any { + value = queue.SanitizeErrorMessage(value, 1024) + if value == "upstream request failed" { + return nil + } + return value +} + +func newRuntimeToken() (string, error) { + value := make([]byte, 16) + if _, err := rand.Read(value); err != nil { + return "", fmt.Errorf("generate route runtime token: %w", err) + } + hexValue := hex.EncodeToString(value) + return hexValue[0:8] + "-" + hexValue[8:12] + "-" + hexValue[12:16] + "-" + hexValue[16:20] + "-" + hexValue[20:32], nil +} + +var _ router.RuntimeRepository = (*GORMRuntime)(nil) diff --git a/portal/worker/worker.go b/portal/worker/worker.go index 14f7300..84f0f4e 100644 --- a/portal/worker/worker.go +++ b/portal/worker/worker.go @@ -3,6 +3,7 @@ package worker import ( "bytes" "context" + "encoding/json" "errors" "fmt" "io" @@ -12,17 +13,19 @@ import ( "git.ilapage.cn/OPC/chorus/internal/core/model" "git.ilapage.cn/OPC/chorus/internal/core/provider" "git.ilapage.cn/OPC/chorus/internal/core/queue" + "git.ilapage.cn/OPC/chorus/internal/core/router" corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage" ) var ( ErrInvalidConfig = errors.New("worker configuration is invalid") - ErrNoProvider = errors.New("exactly one enabled provider model is required") + ErrNoProvider = errors.New("routed provider model is unavailable") ) type Queue interface { ClaimNext(ctx context.Context, owner string, leaseDuration time.Duration) (*queue.Claim, error) - AssignProvider(ctx context.Context, generationID uint64, leaseToken string, providerModelID uint64) (bool, error) + BeginProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, routeMemberID, providerModelID uint64) (model.Attempt, bool, error) + FinishProviderAttempt(ctx context.Context, generationID uint64, leaseToken string, attempt model.Attempt) (bool, error) Succeed(ctx context.Context, generationID uint64, leaseToken string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) Fail(ctx context.Context, generationID uint64, leaseToken string, code, message string, attempt model.Attempt) (bool, error) Inputs(ctx context.Context, generationID uint64) ([]model.GenerationInput, error) @@ -33,6 +36,7 @@ type ClaimController interface{ StopClaims() } type Selection struct { ProviderModelID uint64 BaseURL string + AuthType provider.AuthType APIKey string ModelID string APIType model.APIType @@ -41,7 +45,7 @@ type Selection struct { } type Catalog interface { - Select(ctx context.Context, kind model.GenerationKind) (Selection, error) + Resolve(ctx context.Context, member router.MemberSnapshot, capability model.Capability) (Selection, error) } type ClientFactory interface { New(selection Selection) (provider.Client, error) @@ -65,6 +69,7 @@ type Config struct { Owner string LeaseDuration time.Duration PollInterval time.Duration + Random router.RandomSource } type Worker struct { @@ -73,14 +78,18 @@ type Worker struct { controller ClaimController catalog Catalog factory ClientFactory + runtime router.RuntimeRepository storage ImageStore } -func New(config Config, queueRepository Queue, controller ClaimController, catalog Catalog, factory ClientFactory, storage ImageStore) (*Worker, error) { - if strings.TrimSpace(config.Owner) == "" || config.LeaseDuration <= 0 || config.PollInterval <= 0 || queueRepository == nil || controller == nil || catalog == nil || factory == nil || storage == nil { +func New(config Config, queueRepository Queue, controller ClaimController, catalog Catalog, factory ClientFactory, runtime router.RuntimeRepository, storage ImageStore) (*Worker, error) { + if strings.TrimSpace(config.Owner) == "" || config.LeaseDuration <= 0 || config.PollInterval <= 0 || queueRepository == nil || controller == nil || catalog == nil || factory == nil || runtime == nil || storage == nil { return nil, ErrInvalidConfig } - return &Worker{config: config, queue: queueRepository, controller: controller, catalog: catalog, factory: factory, storage: storage}, nil + if config.Random == nil { + config.Random = router.CryptoRandom{} + } + return &Worker{config: config, queue: queueRepository, controller: controller, catalog: catalog, factory: factory, runtime: runtime, storage: storage}, nil } func (w *Worker) Run(ctx context.Context) error { @@ -130,54 +139,189 @@ func (w *Worker) ProcessOne(ctx context.Context) (bool, error) { } func (w *Worker) process(ctx context.Context, claim *queue.Claim) error { - started := time.Now() - selection, err := w.catalog.Select(ctx, claim.Generation.Kind) + snapshot, err := decodeSnapshot(claim.Generation.RouteSnapshot) if err != nil { - return fmt.Errorf("select worker provider: %w", err) - } - assigned, err := w.queue.AssignProvider(ctx, claim.Generation.ID, claim.LeaseToken, selection.ProviderModelID) - if err != nil || !assigned { - return err - } - client, err := w.factory.New(selection) - if err != nil { - return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) + return w.fail(ctx, claim, provider.CodeUnknown) } inputs, err := w.loadInputs(ctx, claim.Generation) if err != nil { - return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) + return w.fail(ctx, claim, provider.CodeUnknown) } - requestCtx := ctx - var cancel context.CancelFunc - if selection.Timeout > 0 { - requestCtx, cancel = context.WithTimeout(ctx, selection.Timeout) - defer cancel() + used := usedRouteMembers(claim.Generation.Attempts) + if claim.Generation.ProviderAttemptCount >= uint32(snapshot.MaxFailover)+1 { + return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted)) } - generated, err := client.Generate(requestCtx, provider.Request{Kind: claim.Generation.Kind, APIType: selection.APIType, ModelID: selection.ModelID, RenderedPrompt: claim.Generation.RenderedPrompt, Inputs: inputs}) + states, err := w.runtime.MemberStates(ctx, snapshot.Capability, snapshot.Members) if err != nil { - var providerErr *provider.Error - if errors.As(err, &providerErr) { - return w.fail(ctx, claim, selection.ProviderModelID, started, providerErr.Code, providerErr) + return w.fail(ctx, claim, provider.CodeUnknown) + } + for index := range states { + if used[states[index].RoutePoolMemberID] { + states[index].Enabled = false } - return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) } - outputs, keys, err := w.saveOutputs(ctx, claim.Generation, generated) + candidates, err := router.SelectCandidates(snapshot, states, w.config.Random) if err != nil { - w.deleteKeys(keys) - return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) + if errors.Is(err, router.ErrRouteUnavailable) { + if claim.Generation.ProviderAttemptCount > 0 { + return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted)) + } + return w.fail(ctx, claim, provider.ErrorCode(router.CodeRouteUnavailable)) + } + return w.fail(ctx, claim, provider.CodeUnknown) } - attempt := model.Attempt{LatencyMS: elapsedMilliseconds(started)} - finalizeCtx, finalizeCancel := context.WithTimeout(context.Background(), 5*time.Second) - defer finalizeCancel() - owned, err := w.queue.Succeed(finalizeCtx, claim.Generation.ID, claim.LeaseToken, outputs, attempt) - if err != nil || !owned { - w.deleteKeys(keys) + remaining := int(uint32(snapshot.MaxFailover) + 1 - claim.Generation.ProviderAttemptCount) + if remaining < len(candidates) { + candidates = candidates[:remaining] + } + for _, member := range candidates { + reservation, reserveErr := w.runtime.Reserve(ctx, router.ReservationRequest{ + RoutePoolMemberID: member.RoutePoolMemberID, Capability: snapshot.Capability, + Owner: w.config.Owner, LeaseDuration: w.config.LeaseDuration, + }) + if errors.Is(reserveErr, router.ErrMemberUnavailable) { + continue + } + if reserveErr != nil { + return w.fail(ctx, claim, provider.CodeUnknown) + } + selection, resolveErr := w.catalog.Resolve(ctx, member, snapshot.Capability) + if resolveErr != nil { + continue + } + attempt, begun, beginErr := w.queue.BeginProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, member.RoutePoolMemberID, selection.ProviderModelID) + if beginErr != nil { + return beginErr + } + if !begun { + return nil + } + client, factoryErr := w.factory.New(selection) + if factoryErr != nil { + finished, err := w.finishFailure(ctx, claim, reservation, attempt, provider.CodeUnknown, provider.FailureOther) + if err != nil { + return err + } + if !finished { + return nil + } + return w.fail(ctx, claim, provider.CodeUnknown) + } + requestCtx := ctx + var cancel context.CancelFunc + if selection.Timeout > 0 { + requestCtx, cancel = context.WithTimeout(ctx, selection.Timeout) + } + generated, generateErr := client.Generate(requestCtx, provider.Request{ + Kind: claim.Generation.Kind, APIType: selection.APIType, ModelID: selection.ModelID, + RenderedPrompt: claim.Generation.RenderedPrompt, Inputs: inputs, + }) + if cancel != nil { + cancel() + } + if generateErr != nil { + code, class := provider.CodeUnknown, provider.FailureOther + var providerErr *provider.Error + if errors.As(generateErr, &providerErr) { + code, class = providerErr.Code, providerErr.Class + } + finished, err := w.finishFailure(ctx, claim, reservation, attempt, code, class) + if err != nil { + return err + } + if !finished { + return nil + } + if router.ActionForFailure(class) == router.FailureTryNext { + continue + } + return w.fail(ctx, claim, code) + } + finished, err := w.finishSuccess(ctx, claim, reservation, attempt) if err != nil { return err } + if !finished { + return nil + } + outputs, keys, saveErr := w.saveOutputs(ctx, claim.Generation, claim.LeaseToken, attempt.ProviderOrdinal, generated) + if saveErr != nil { + w.deleteKeys(keys) + return w.fail(ctx, claim, provider.CodeUnknown) + } + finalizeCtx, finalizeCancel := context.WithTimeout(context.Background(), 5*time.Second) + owned, succeedErr := w.queue.Succeed(finalizeCtx, claim.Generation.ID, claim.LeaseToken, outputs, model.Attempt{}) + finalizeCancel() + if succeedErr != nil || !owned { + w.deleteKeys(keys) + return succeedErr + } return nil } - return nil + if claim.Generation.ProviderAttemptCount > 0 { + return w.fail(ctx, claim, provider.ErrorCode(router.CodeFailoverExhausted)) + } + return w.fail(ctx, claim, provider.ErrorCode(router.CodeRouteUnavailable)) +} + +func decodeSnapshot(encoded []byte) (router.RouteSnapshot, error) { + var snapshot router.RouteSnapshot + if len(encoded) == 0 || json.Unmarshal(encoded, &snapshot) != nil { + return router.RouteSnapshot{}, router.ErrInvalidSnapshot + } + if err := snapshot.Validate(); err != nil { + return router.RouteSnapshot{}, err + } + return snapshot, nil +} + +func usedRouteMembers(encoded []byte) map[uint64]bool { + var attempts []model.Attempt + if json.Unmarshal(encoded, &attempts) != nil { + return map[uint64]bool{} + } + used := make(map[uint64]bool, len(attempts)) + for _, attempt := range attempts { + if attempt.Type == "provider" && attempt.RouteMemberID != 0 { + used[attempt.RouteMemberID] = true + } + } + return used +} + +func (w *Worker) finishFailure(ctx context.Context, claim *queue.Claim, reservation router.Reservation, attempt model.Attempt, code provider.ErrorCode, class provider.FailureClass) (bool, error) { + retryable := router.Retryable(class) + attempt.ErrorCode = string(code) + attempt.ErrorMessage = "upstream request failed" + attempt.LatencyMS = elapsedMilliseconds(attempt.StartedAt) + attempt.Retryable = &retryable + finished, err := w.queue.FinishProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, attempt) + if err != nil { + return false, err + } + _, recordErr := w.runtime.Record(ctx, reservation, router.CircuitObservation{ + Retryable: retryable, OpenImmediately: class == provider.FailureUnauthorized, + ErrorCode: string(code), ErrorMessage: "upstream request failed", + }) + if recordErr != nil { + return false, recordErr + } + return finished, nil +} + +func (w *Worker) finishSuccess(ctx context.Context, claim *queue.Claim, reservation router.Reservation, attempt model.Attempt) (bool, error) { + retryable := false + attempt.LatencyMS = elapsedMilliseconds(attempt.StartedAt) + attempt.Retryable = &retryable + finished, err := w.queue.FinishProviderAttempt(ctx, claim.Generation.ID, claim.LeaseToken, attempt) + if err != nil { + return false, err + } + _, recordErr := w.runtime.Record(ctx, reservation, router.CircuitObservation{Succeeded: true}) + if recordErr != nil { + return false, recordErr + } + return finished, nil } func (w *Worker) loadInputs(ctx context.Context, generation model.Generation) ([]provider.Input, error) { @@ -207,7 +351,7 @@ func (w *Worker) loadInputs(ctx context.Context, generation model.Generation) ([ return inputs, nil } -func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, generated []provider.Output) ([]model.GenerationOutput, []string, error) { +func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, leaseToken string, providerOrdinal uint32, generated []provider.Output) ([]model.GenerationOutput, []string, error) { if len(generated) == 0 { return nil, nil, fmt.Errorf("provider returned no outputs") } @@ -219,7 +363,10 @@ func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, g rows = append(rows, model.GenerationOutput{Kind: output.Kind, TextContent: &text}) continue } - key := fmt.Sprintf("outputs/%d/%d/%d-%d", generation.UserID, generation.ID, generation.AttemptCount, index+1) + if output.Kind != model.KindImage || len(output.Content) == 0 { + return nil, keys, fmt.Errorf("provider returned invalid output") + } + key := fmt.Sprintf("outputs/%d/%d/%s/%d-%d", generation.UserID, generation.ID, leaseToken, providerOrdinal, index+1) thumb := key + "-thumbnail" objects, err := w.storage.PutImage(ctx, ImageRequest{Key: key, ThumbnailKey: thumb, OwnerID: generation.UserID, GenerationID: generation.ID, ContentType: output.ContentType, Source: bytes.NewReader(output.Content)}) if err != nil { @@ -234,25 +381,24 @@ func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, g return rows, keys, nil } -func (w *Worker) fail(ctx context.Context, claim *queue.Claim, providerModelID uint64, started time.Time, code provider.ErrorCode, cause error) error { - message := "upstream request failed" - attempt := model.Attempt{ProviderModelID: providerModelID, LatencyMS: elapsedMilliseconds(started), ErrorCode: string(code), ErrorMessage: message} +func (w *Worker) fail(ctx context.Context, claim *queue.Claim, code provider.ErrorCode) error { finalizeCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - _, err := w.queue.Fail(finalizeCtx, claim.Generation.ID, claim.LeaseToken, string(code), message, attempt) - if err != nil { - return err - } - return nil + _, err := w.queue.Fail(finalizeCtx, claim.Generation.ID, claim.LeaseToken, string(code), "upstream request failed", model.Attempt{}) + return err } + func (w *Worker) deleteKeys(keys []string) { for _, key := range keys { _ = w.storage.Delete(context.Background(), key) } } -func elapsedMilliseconds(started time.Time) int64 { - elapsed := time.Since(started).Milliseconds() +func elapsedMilliseconds(started *time.Time) int64 { + if started == nil { + return 1 + } + elapsed := time.Since(*started).Milliseconds() if elapsed < 1 { return 1 } diff --git a/portal/worker/worker_test.go b/portal/worker/worker_test.go index 81cd3c3..9d9ddad 100644 --- a/portal/worker/worker_test.go +++ b/portal/worker/worker_test.go @@ -3,229 +3,241 @@ package worker import ( "bytes" "context" + "errors" "image" "image/color" "image/png" - "sync" + "io" "testing" "time" "git.ilapage.cn/OPC/chorus/internal/core/model" "git.ilapage.cn/OPC/chorus/internal/core/provider" "git.ilapage.cn/OPC/chorus/internal/core/queue" + "git.ilapage.cn/OPC/chorus/internal/core/router" corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage" - platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage" ) type fakeQueue struct { - mutex sync.Mutex - claim *queue.Claim - inputs []model.GenerationInput - succeeded []model.GenerationOutput - failedCode string - attempt model.Attempt - stale bool + claim *queue.Claim + inputs []model.GenerationInput + begun []model.Attempt + finished []model.Attempt + succeeded []model.GenerationOutput + failedCode string + staleSucceed bool } func (q *fakeQueue) ClaimNext(context.Context, string, time.Duration) (*queue.Claim, error) { - q.mutex.Lock() - defer q.mutex.Unlock() claim := q.claim q.claim = nil return claim, nil } -func (q *fakeQueue) AssignProvider(context.Context, uint64, string, uint64) (bool, error) { - return !q.stale, nil +func (q *fakeQueue) BeginProviderAttempt(_ context.Context, _ uint64, _ string, memberID, providerModelID uint64) (model.Attempt, bool, error) { + started := time.Now() + attempt := model.Attempt{Type: "provider", RouteMemberID: memberID, ProviderModelID: providerModelID, ProviderOrdinal: uint32(len(q.begun) + 1), StartedAt: &started} + q.begun = append(q.begun, attempt) + return attempt, true, nil } -func (q *fakeQueue) Succeed(_ context.Context, _ uint64, _ string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) { - q.succeeded = outputs - q.attempt = attempt - return !q.stale, nil +func (q *fakeQueue) FinishProviderAttempt(_ context.Context, _ uint64, _ string, attempt model.Attempt) (bool, error) { + q.finished = append(q.finished, attempt) + return true, nil } -func (q *fakeQueue) Fail(_ context.Context, _ uint64, _ string, code, _ string, attempt model.Attempt) (bool, error) { +func (q *fakeQueue) Succeed(_ context.Context, _ uint64, _ string, outputs []model.GenerationOutput, _ model.Attempt) (bool, error) { + if q.staleSucceed { + return false, nil + } + q.succeeded = append(q.succeeded, outputs...) + return true, nil +} +func (q *fakeQueue) Fail(_ context.Context, _ uint64, _ string, code, _ string, _ model.Attempt) (bool, error) { q.failedCode = code - q.attempt = attempt - return !q.stale, nil + return true, nil } func (q *fakeQueue) Inputs(context.Context, uint64) ([]model.GenerationInput, error) { return q.inputs, nil } -type fakeController struct { - mutex sync.Mutex - stopped bool +type fakeController struct{ stopped bool } + +func (c *fakeController) StopClaims() { c.stopped = true } + +type fakeRuntime struct { + states []router.MemberState + unavailable map[uint64]bool + reservations []router.ReservationRequest + observations []router.CircuitObservation } -func (c *fakeController) StopClaims() { c.mutex.Lock(); c.stopped = true; c.mutex.Unlock() } - -type fakeCatalog struct{ selection Selection } - -func (c fakeCatalog) Select(context.Context, model.GenerationKind) (Selection, error) { - return c.selection, nil +func (r *fakeRuntime) MemberStates(context.Context, model.Capability, []router.MemberSnapshot) ([]router.MemberState, error) { + return append([]router.MemberState(nil), r.states...), nil +} +func (r *fakeRuntime) Reserve(_ context.Context, request router.ReservationRequest) (router.Reservation, error) { + r.reservations = append(r.reservations, request) + if r.unavailable[request.RoutePoolMemberID] { + return router.Reservation{}, router.ErrMemberUnavailable + } + return router.Reservation{RoutePoolMemberID: request.RoutePoolMemberID, State: router.CircuitClosed}, nil +} +func (r *fakeRuntime) Record(_ context.Context, _ router.Reservation, observation router.CircuitObservation) (bool, error) { + r.observations = append(r.observations, observation) + return true, nil } -type fakeFactory struct{ client provider.Client } +type fakeCatalog struct{ selections map[uint64]Selection } -func (f fakeFactory) New(Selection) (provider.Client, error) { return f.client, nil } +func (c fakeCatalog) Resolve(_ context.Context, member router.MemberSnapshot, _ model.Capability) (Selection, error) { + selection, ok := c.selections[member.RoutePoolMemberID] + if !ok { + return Selection{}, ErrNoProvider + } + return selection, nil +} + +type fakeFactory struct { + clients []provider.Client + index int +} + +func (f *fakeFactory) New(Selection) (provider.Client, error) { + if f.index >= len(f.clients) { + return nil, errors.New("no fake provider") + } + client := f.clients[f.index] + f.index++ + return client, nil +} type fakeProvider struct { outputs []provider.Output err error - started chan struct{} - release chan struct{} } -func (p *fakeProvider) Generate(ctx context.Context, _ provider.Request) ([]provider.Output, error) { - if p.started != nil { - close(p.started) - } - if p.release != nil { - select { - case <-p.release: - case <-ctx.Done(): - return nil, ctx.Err() - } - } +func (p fakeProvider) Generate(context.Context, provider.Request) ([]provider.Output, error) { return p.outputs, p.err } -func TestTextAndImageSuccessPersistResults(t *testing.T) { - t.Run("text", func(t *testing.T) { - q := newQueue(model.KindText) - store := newStore(t) - w := newWorker(t, q, &fakeController{}, fakeCatalog{selection()}, fakeFactory{&fakeProvider{outputs: []provider.Output{{Kind: model.KindText, Text: "done"}}}}, store) - worked, err := w.ProcessOne(context.Background()) - if err != nil || !worked || len(q.succeeded) != 1 || q.succeeded[0].TextContent == nil || *q.succeeded[0].TextContent != "done" { - t.Fatalf("result=%v err=%v rows=%#v", worked, err, q.succeeded) - } - }) - t.Run("image", func(t *testing.T) { - q := newQueue(model.KindImage) - store := newStore(t) - data := testPNG() - _, err := store.storage.Put(context.Background(), corestorage.PutRequest{Key: "inputs/1", OwnerID: 1, GenerationID: 10, ContentType: "image/png", Source: bytes.NewReader(data)}) - if err != nil { - t.Fatal(err) - } - q.inputs = []model.GenerationInput{{GenerationID: 10, Position: 0, Role: model.RolePrimary, MIMEType: "image/png", StorageKey: "inputs/1"}} - sel := selection() - sel.APIType = model.APIImagesEdits - w := newWorker(t, q, &fakeController{}, fakeCatalog{sel}, fakeFactory{&fakeProvider{outputs: []provider.Output{{Kind: model.KindImage, Content: data, ContentType: "image/png"}}}}, store) - worked, err := w.ProcessOne(context.Background()) - if err != nil || !worked || len(q.succeeded) != 1 || q.succeeded[0].StorageKey == nil || q.succeeded[0].ThumbnailStorageKey == nil { - t.Fatalf("result=%v err=%v rows=%#v", worked, err, q.succeeded) - } - for _, key := range []string{*q.succeeded[0].StorageKey, *q.succeeded[0].ThumbnailStorageKey} { - reader, _, err := store.Open(context.Background(), key) - if err != nil { - t.Fatal(err) - } - reader.Close() - } - }) +type fakeStore struct { + deleted []string } -func TestStaleTokenDeletesGeneratedFiles(t *testing.T) { - q := newQueue(model.KindImage) - q.stale = false - store := newStore(t) - data := testPNG() - _, _ = store.storage.Put(context.Background(), corestorage.PutRequest{Key: "inputs/1", OwnerID: 1, GenerationID: 10, ContentType: "image/png", Source: bytes.NewReader(data)}) - q.inputs = []model.GenerationInput{{GenerationID: 10, MIMEType: "image/png", StorageKey: "inputs/1"}} - client := &fakeProvider{outputs: []provider.Output{{Kind: model.KindImage, Content: data, ContentType: "image/png"}}} - wrapped := &staleOnSucceedQueue{fakeQueue: q} - sel := selection() - sel.APIType = model.APIImagesEdits - w := newWorker(t, wrapped, &fakeController{}, fakeCatalog{sel}, fakeFactory{client}, store) +func (s *fakeStore) Open(context.Context, string) (io.ReadCloser, corestorage.Object, error) { + return nil, corestorage.Object{}, errors.New("unexpected input read") +} +func (s *fakeStore) PutImage(_ context.Context, request ImageRequest) (ImageObjects, error) { + return ImageObjects{ + Original: corestorage.Object{Key: request.Key, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: request.ContentType, Size: 1}, + Thumbnail: corestorage.Object{Key: request.ThumbnailKey, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: "image/png", Size: 1}, + }, nil +} +func (s *fakeStore) Delete(_ context.Context, key string) error { + s.deleted = append(s.deleted, key) + return nil +} + +type firstRandom struct{} + +func (firstRandom) Uint64n(uint64) uint64 { return 0 } + +func TestWorkerRetriesOnlyRetryableProviderFailure(t *testing.T) { + q := newQueue(t, model.KindText, model.CapabilityText, 1) + runtime := &fakeRuntime{states: activeStates(1, 2)} + factory := &fakeFactory{clients: []provider.Client{ + fakeProvider{err: &provider.Error{Code: provider.CodeRateLimited, Class: provider.FailureRateLimited}}, + fakeProvider{outputs: []provider.Output{{Kind: model.KindText, Text: "done"}}}, + }} + w := newWorker(t, q, runtime, fakeCatalog{selections: selections()}, factory, &fakeStore{}) + worked, err := w.ProcessOne(context.Background()) + if err != nil || !worked || len(q.begun) != 2 || len(q.finished) != 2 || len(q.succeeded) != 1 || q.failedCode != "" { + t.Fatalf("worked=%v err=%v begun=%#v finished=%#v succeeded=%#v failed=%s", worked, err, q.begun, q.finished, q.succeeded, q.failedCode) + } + if len(runtime.observations) != 2 || !runtime.observations[0].Retryable || !runtime.observations[1].Succeeded { + t.Fatalf("runtime observations=%#v", runtime.observations) + } +} + +func TestWorkerStopsOnUnauthorizedAndOpensCircuit(t *testing.T) { + q := newQueue(t, model.KindText, model.CapabilityText, 1) + runtime := &fakeRuntime{states: activeStates(1, 2)} + factory := &fakeFactory{clients: []provider.Client{fakeProvider{err: &provider.Error{Code: provider.CodeUnauthorized, Class: provider.FailureUnauthorized}}}} + w := newWorker(t, q, runtime, fakeCatalog{selections: selections()}, factory, &fakeStore{}) _, err := w.ProcessOne(context.Background()) - if err != nil { - t.Fatal(err) + if err != nil || len(q.begun) != 1 || q.failedCode != string(provider.CodeUnauthorized) { + t.Fatalf("err=%v begun=%#v failed=%s", err, q.begun, q.failedCode) } - if _, _, err := store.Open(context.Background(), "outputs/1/10/1-1"); err == nil { - t.Fatal("stale output was retained") + if len(runtime.observations) != 1 || !runtime.observations[0].OpenImmediately || runtime.observations[0].Retryable { + t.Fatalf("runtime observations=%#v", runtime.observations) } } -type staleOnSucceedQueue struct{ *fakeQueue } - -func (q *staleOnSucceedQueue) Succeed(ctx context.Context, id uint64, token string, outputs []model.GenerationOutput, attempt model.Attempt) (bool, error) { - q.fakeQueue.Succeed(ctx, id, token, outputs, attempt) - return false, nil -} - -func TestProviderFailureIsPersistedWithoutSensitiveDetail(t *testing.T) { - q := newQueue(model.KindText) - store := newStore(t) - client := &fakeProvider{err: &provider.Error{Code: provider.CodeRateLimited, Class: provider.FailureRateLimited, Message: "secret response"}} - w := newWorker(t, q, &fakeController{}, fakeCatalog{selection()}, fakeFactory{client}, store) +func TestWorkerCleansStagedOutputAfterStaleSucceed(t *testing.T) { + q := newQueue(t, model.KindImage, model.CapabilityImageGenerate, 0) + q.staleSucceed = true + runtime := &fakeRuntime{states: activeStates(1)} + store := &fakeStore{} + w := newWorker(t, q, runtime, fakeCatalog{selections: map[uint64]Selection{1: {ProviderModelID: 101, APIType: model.APIImages, ModelID: "image"}}}, &fakeFactory{clients: []provider.Client{fakeProvider{outputs: []provider.Output{{Kind: model.KindImage, Content: []byte("image"), ContentType: "image/png"}}}}}, store) _, err := w.ProcessOne(context.Background()) - if err != nil { - t.Fatal(err) - } - if q.failedCode != string(provider.CodeRateLimited) || q.attempt.ErrorMessage != "upstream request failed" { - t.Fatalf("failure=%q attempt=%#v", q.failedCode, q.attempt) + if err != nil || len(q.succeeded) != 0 || len(store.deleted) != 2 { + t.Fatalf("err=%v succeeded=%#v deleted=%#v", err, q.succeeded, store.deleted) } } -func TestRunStopsClaimsAndWaitsForInflight(t *testing.T) { - q := newQueue(model.KindText) - controller := &fakeController{} - client := &fakeProvider{outputs: []provider.Output{{Kind: model.KindText, Text: "done"}}, started: make(chan struct{}), release: make(chan struct{})} - w := newWorker(t, q, controller, fakeCatalog{selection()}, fakeFactory{client}, newStore(t)) - ctx, cancel := context.WithCancel(context.Background()) - done := make(chan error, 1) - go func() { done <- w.Run(ctx) }() - <-client.started - cancel() - time.Sleep(20 * time.Millisecond) - controller.mutex.Lock() - stopped := controller.stopped - controller.mutex.Unlock() - if !stopped { - t.Fatal("claims were not stopped") - } - select { - case <-done: - t.Fatal("worker returned before inflight completion") - default: - } - close(client.release) - select { - case err := <-done: - if err != nil { - t.Fatal(err) - } - case <-time.After(time.Second): - t.Fatal("worker did not exit") +func TestWorkerFailsWhenRouteSnapshotIsMissing(t *testing.T) { + q := &fakeQueue{claim: &queue.Claim{Generation: model.Generation{ID: 10, UserID: 1, Kind: model.KindText, Attempts: []byte("[]")}, LeaseToken: "token"}} + w := newWorker(t, q, &fakeRuntime{}, fakeCatalog{}, &fakeFactory{}, &fakeStore{}) + _, err := w.ProcessOne(context.Background()) + if err != nil || q.failedCode != string(provider.CodeUnknown) { + t.Fatalf("err=%v failed=%s", err, q.failedCode) } } -func newQueue(kind model.GenerationKind) *fakeQueue { - return &fakeQueue{claim: &queue.Claim{Generation: model.Generation{ID: 10, UserID: 1, Kind: kind, RenderedPrompt: "prompt", AttemptCount: 1}, LeaseToken: "token"}} -} -func selection() Selection { - return Selection{ProviderModelID: 2, BaseURL: "https://provider.test/v1", ModelID: "mock", APIType: model.APIChat, Timeout: time.Second} -} -func newWorker(t *testing.T, q Queue, c ClaimController, catalog Catalog, factory ClientFactory, store ImageStore) *Worker { +func newQueue(t *testing.T, kind model.GenerationKind, capability model.Capability, maxFailover uint16) *fakeQueue { t.Helper() - w, err := New(Config{Owner: "test", LeaseDuration: time.Second, PollInterval: time.Millisecond}, q, c, catalog, factory, store) + snapshot := router.RouteSnapshot{ + Capability: capability, RoutePoolID: 1, RoutePoolVersion: 1, PromptTemplateID: 1, + PromptTemplateKey: "test", PromptTemplateVersion: 1, MaxFailover: maxFailover, + Members: []router.MemberSnapshot{ + {RoutePoolMemberID: 1, ProviderModelID: 101, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}, + {RoutePoolMemberID: 2, ProviderModelID: 102, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}, + }, + } + if maxFailover == 0 { + snapshot.Members = snapshot.Members[:1] + } + encoded, err := snapshot.Encode() + if err != nil { + t.Fatal(err) + } + return &fakeQueue{claim: &queue.Claim{Generation: model.Generation{ + ID: 10, UserID: 1, Kind: kind, RenderedPrompt: "prompt", RouteSnapshot: encoded, + Attempts: []byte("[]"), AttemptCount: 1, + }, LeaseToken: "token"}} +} + +func activeStates(memberIDs ...uint64) []router.MemberState { + states := make([]router.MemberState, 0, len(memberIDs)) + for _, memberID := range memberIDs { + states = append(states, router.MemberState{RoutePoolMemberID: memberID, Enabled: true, ProviderEnabled: true, ModelEnabled: true, SupportsCapability: true, CircuitState: router.CircuitClosed, HalfOpenMax: 1}) + } + return states +} + +func selections() map[uint64]Selection { + return map[uint64]Selection{ + 1: {ProviderModelID: 101, APIType: model.APIChat, ModelID: "first"}, + 2: {ProviderModelID: 102, APIType: model.APIChat, ModelID: "second"}, + } +} + +func newWorker(t *testing.T, q Queue, runtime router.RuntimeRepository, catalog Catalog, factory ClientFactory, store ImageStore) *Worker { + t.Helper() + w, err := New(Config{Owner: "test", LeaseDuration: time.Second, PollInterval: time.Millisecond, Random: firstRandom{}}, q, &fakeController{}, catalog, factory, runtime, store) if err != nil { t.Fatal(err) } return w } -func newStore(t *testing.T) *LocalStorage { - t.Helper() - local, err := platformstorage.NewLocal(platformstorage.Config{Root: t.TempDir(), MaxObjectBytes: 1 << 20, MaxImagePixels: 10000, ThumbnailMaxSide: 32, AllowedImageMIME: map[string]bool{"image/png": true}}) - if err != nil { - t.Fatal(err) - } - store, err := NewLocalStorage(local) - if err != nil { - t.Fatal(err) - } - return store -} + func testPNG() []byte { img := image.NewRGBA(image.Rect(0, 0, 4, 4)) for y := 0; y < 4; y++ {