feat: 实现安全多 Provider worker (#22)

This commit is contained in:
ila
2026-08-21 21:00:25 +08:00
parent 9b2827a727
commit 35fd0342f0
22 changed files with 1498 additions and 375 deletions
+4
View File
@@ -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"`
+179 -23
View File
@@ -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"}
}
+1 -1
View File
@@ -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)
}
+97
View File
@@ -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)
}
}
}
+14
View File
@@ -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)
}
+8
View File
@@ -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)
}
+6
View File
@@ -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
}
@@ -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)
}
}
+102 -10
View File
@@ -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)
+2
View File
@@ -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)
+7
View File
@@ -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 ""
+24
View File
@@ -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{}
+8 -3
View File
@@ -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
}
+65 -8
View File
@@ -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"):
@@ -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())
}
}
+8
View File
@@ -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)
+5 -1
View File
@@ -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
}
+52 -25
View File
@@ -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,
})
}
+194 -93
View File
@@ -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)
}
}
+289
View File
@@ -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)
+196 -50
View File
@@ -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
}
+173 -161
View File
@@ -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++ {