feat: 实现安全多 Provider worker (#22)
This commit is contained in:
@@ -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"`
|
||||
|
||||
@@ -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"}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
@@ -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{}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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++ {
|
||||
|
||||
Reference in New Issue
Block a user