diff --git a/cmd/chorus-mock-provider/main.go b/cmd/chorus-mock-provider/main.go new file mode 100644 index 0000000..01acfbb --- /dev/null +++ b/cmd/chorus-mock-provider/main.go @@ -0,0 +1,22 @@ +package main + +import ( + "log" + "net/http" + "os" + "time" + + "git.ilapage.cn/OPC/chorus/internal/platform/mockprovider" +) + +func main() { + address := os.Getenv("CHORUS_MOCK_ADDR") + if address == "" { + address = "127.0.0.1:18080" + } + server := &http.Server{Addr: address, Handler: mockprovider.Handler{}, ReadHeaderTimeout: 5 * time.Second} + log.Printf("mock provider listening on %s", address) + if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatal(err) + } +} diff --git a/internal/core/provider/openai.go b/internal/core/provider/openai.go new file mode 100644 index 0000000..9ab50f1 --- /dev/null +++ b/internal/core/provider/openai.go @@ -0,0 +1,264 @@ +package provider + +import ( + "bytes" + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "mime/multipart" + "net/http" + "net/url" + "path" + "strings" +) + +const defaultMaxResponseBytes int64 = 20 << 20 + +var ( + ErrInvalidConfig = errors.New("provider configuration is invalid") + ErrInvalidRequest = errors.New("provider request is invalid") +) + +type OpenAIConfig struct { + BaseURL string + APIKey string + ExtraBody json.RawMessage + MaxResponseBytes int64 +} + +type OpenAI struct { + http HTTPClient + baseURL *url.URL + apiKey string + extraBody map[string]any + maxResponseBytes int64 +} + +func NewOpenAI(httpClient HTTPClient, config OpenAIConfig) (*OpenAI, error) { + if httpClient == nil || strings.TrimSpace(config.BaseURL) == "" { + return nil, ErrInvalidConfig + } + baseURL, err := url.Parse(config.BaseURL) + if err != nil || (baseURL.Scheme != "http" && baseURL.Scheme != "https") || baseURL.Hostname() == "" || baseURL.User != nil || baseURL.RawQuery != "" || baseURL.Fragment != "" { + return nil, ErrInvalidConfig + } + extraBody, err := parseExtraBody(config.ExtraBody) + if err != nil { + return nil, err + } + if config.MaxResponseBytes <= 0 { + config.MaxResponseBytes = defaultMaxResponseBytes + } + return &OpenAI{http: httpClient, baseURL: baseURL, apiKey: config.APIKey, extraBody: extraBody, maxResponseBytes: config.MaxResponseBytes}, nil +} + +func (c *OpenAI) Generate(ctx context.Context, request Request) ([]Output, error) { + if strings.TrimSpace(request.ModelID) == "" || strings.TrimSpace(request.RenderedPrompt) == "" { + return nil, ErrInvalidRequest + } + switch request.APIType { + case "chat": + return c.chat(ctx, request) + case "images_edits": + return c.imagesEdits(ctx, request) + default: + return nil, ErrInvalidRequest + } +} + +func (c *OpenAI) chat(ctx context.Context, request Request) ([]Output, error) { + body := map[string]any{ + "model": request.ModelID, + "messages": []map[string]string{{"role": "user", "content": request.RenderedPrompt}}, + } + mergeExtra(body, c.extraBody) + encoded, err := json.Marshal(body) + if err != nil { + return nil, ErrInvalidRequest + } + response, err := c.do(ctx, "/chat/completions", "application/json", encoded) + if err != nil { + return nil, err + } + var decoded struct { + Choices []struct { + Message struct { + Content string `json:"content"` + } `json:"message"` + } `json:"choices"` + } + if json.Unmarshal(response.Body, &decoded) != nil || len(decoded.Choices) == 0 || decoded.Choices[0].Message.Content == "" { + return nil, providerError(CodeUnknown, FailureOther) + } + 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 { + return nil, ErrInvalidRequest + } + var body bytes.Buffer + writer := multipart.NewWriter(&body) + if err := writer.WriteField("model", request.ModelID); err != nil { + return nil, ErrInvalidRequest + } + if err := writer.WriteField("prompt", request.RenderedPrompt); err != nil { + return nil, ErrInvalidRequest + } + for key, value := range c.extraBody { + encoded, err := scalarString(value) + if err != nil || writer.WriteField(key, encoded) != nil { + return nil, ErrInvalidRequest + } + } + 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 + } + if _, err := part.Write(input.Content); err != nil { + return nil, ErrInvalidRequest + } + } + if err := writer.Close(); err != nil { + return nil, ErrInvalidRequest + } + response, err := c.do(ctx, "/images/edits", writer.FormDataContentType(), body.Bytes()) + if err != nil { + return nil, err + } + 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 { + 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 != "" { + content, err = base64.StdEncoding.DecodeString(item.B64JSON) + if err != nil { + return nil, providerError(CodeUnknown, FailureOther) + } + } else if item.URL != "" { + fetched, fetchErr := c.http.Fetch(ctx, item.URL, c.maxResponseBytes) + if fetchErr != nil { + return nil, classifyNetworkError(fetchErr) + } + if fetched.StatusCode < 200 || fetched.StatusCode >= 300 { + return nil, fromHTTPStatus(fetched.StatusCode, nil) + } + content, contentType = fetched.Body, strings.TrimSpace(strings.Split(fetched.ContentType, ";")[0]) + } else { + return nil, providerError(CodeUnknown, FailureOther) + } + outputs = append(outputs, Output{Kind: request.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) + header := http.Header{"Content-Type": []string{contentType}, "Accept": []string{"application/json"}} + if c.apiKey != "" { + header.Set("Authorization", "Bearer "+c.apiKey) + } + response, err := c.http.Do(ctx, HTTPRequest{Method: http.MethodPost, URL: target.String(), Header: header, Body: body, MaxBytes: c.maxResponseBytes}) + if err != nil { + return HTTPResponse{}, classifyNetworkError(err) + } + if response.StatusCode < 200 || response.StatusCode >= 300 { + return HTTPResponse{}, fromHTTPStatus(response.StatusCode, response.Body) + } + return response, nil +} + +func parseExtraBody(raw json.RawMessage) (map[string]any, error) { + result := map[string]any{} + if len(raw) == 0 || string(raw) == "null" { + return result, nil + } + 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} + for key, value := range result { + if !allowed[key] || key == "model" || key == "messages" || key == "prompt" { + return nil, ErrInvalidConfig + } + switch value.(type) { + case string, float64, bool: + default: + return nil, ErrInvalidConfig + } + } + return result, nil +} + +func mergeExtra(target, extra map[string]any) { + for key, value := range extra { + target[key] = value + } +} +func scalarString(value any) (string, error) { + switch value := value.(type) { + case string: + return value, nil + case float64: + return fmt.Sprintf("%v", value), nil + case bool: + return fmt.Sprintf("%t", value), nil + default: + return "", ErrInvalidConfig + } +} +func imageExtension(mimeType string) string { + if mimeType == "image/jpeg" { + return ".jpg" + } + return ".png" +} + +func fromHTTPStatus(status int, body []byte) *Error { + if status == http.StatusBadRequest && policyRejected(body) { + return providerError(CodePolicyRejected, FailurePolicyRejected) + } + switch status { + case http.StatusTooManyRequests: + return providerError(CodeRateLimited, FailureRateLimited) + case http.StatusBadRequest: + return providerError(CodeBadRequest, FailureBadRequest) + case http.StatusUnauthorized: + return providerError(CodeUnauthorized, FailureUnauthorized) + } + if status >= 500 { + return providerError(CodeServerError, FailureServer) + } + 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"} +} + +var _ Client = (*OpenAI)(nil) diff --git a/internal/core/provider/openai_test.go b/internal/core/provider/openai_test.go new file mode 100644 index 0000000..32154d4 --- /dev/null +++ b/internal/core/provider/openai_test.go @@ -0,0 +1,109 @@ +package provider + +import ( + "context" + "encoding/base64" + "errors" + "net/http" + "strings" + "testing" + + "git.ilapage.cn/OPC/chorus/internal/core/model" +) + +type fakeHTTP struct { + response HTTPResponse + err error + fetched HTTPResponse + request HTTPRequest + fetchURL string +} + +func (f *fakeHTTP) Do(_ context.Context, request HTTPRequest) (HTTPResponse, error) { + f.request = request + return f.response, f.err +} +func (f *fakeHTTP) Fetch(_ context.Context, rawURL string, _ int64) (HTTPResponse, error) { + f.fetchURL = rawURL + return f.fetched, f.err +} + +func TestChatProtocol(t *testing.T) { + httpClient := &fakeHTTP{response: HTTPResponse{StatusCode: 200, ContentType: "application/json", Body: []byte(`{"choices":[{"message":{"content":"done"}}]}`)}} + client, err := NewOpenAI(httpClient, OpenAIConfig{BaseURL: "https://provider.test/v1", APIKey: "secret", ExtraBody: jsonBytes(`{"temperature":0.2}`)}) + if err != nil { + t.Fatal(err) + } + outputs, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "mock-chat", RenderedPrompt: "write"}) + if err != nil || len(outputs) != 1 || outputs[0].Text != "done" { + t.Fatalf("Generate() = %#v, %v", outputs, err) + } + if httpClient.request.URL != "https://provider.test/v1/chat/completions" || httpClient.request.Header.Get("Authorization") != "Bearer secret" || !strings.Contains(string(httpClient.request.Body), `"temperature":0.2`) { + t.Fatalf("request = %#v body=%s", httpClient.request, string(httpClient.request.Body)) + } +} + +func TestImagesEditsBase64AndURL(t *testing.T) { + pngData := []byte("image-bytes") + 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}}}) + 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) + } + if !strings.Contains(httpClient.request.Header.Get("Content-Type"), "multipart/form-data") || !strings.Contains(string(httpClient.request.Body), "mock-image") { + t.Fatalf("multipart request missing fields") + } +} + +func TestFailureClassificationMatrix(t *testing.T) { + tests := []struct { + name string + status int + body string + class FailureClass + code ErrorCode + }{ + {"429", 429, "", FailureRateLimited, CodeRateLimited}, {"500", 500, "", FailureServer, CodeServerError}, + {"400", 400, "", FailureBadRequest, CodeBadRequest}, {"401", 401, "", FailureUnauthorized, CodeUnauthorized}, + {"policy", 400, `{"error":{"code":"content_policy_violation"}}`, FailurePolicyRejected, CodePolicyRejected}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + client, _ := NewOpenAI(&fakeHTTP{response: HTTPResponse{StatusCode: test.status, Body: []byte(test.body)}}, OpenAIConfig{BaseURL: "https://provider.test/v1"}) + _, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"}) + var providerErr *Error + if !errors.As(err, &providerErr) || providerErr.Class != test.class || providerErr.Code != test.code || (test.body != "" && strings.Contains(providerErr.Error(), test.body)) { + t.Fatalf("error = %#v", err) + } + }) + } + for _, test := range []struct { + name string + err error + class FailureClass + }{{"timeout", context.DeadlineExceeded, FailureTimeout}, {"connection", errors.New("dial failed"), FailureConnection}} { + t.Run(test.name, func(t *testing.T) { + client, _ := NewOpenAI(&fakeHTTP{err: test.err}, OpenAIConfig{BaseURL: "https://provider.test/v1"}) + _, err := client.Generate(context.Background(), Request{Kind: model.KindText, APIType: model.APIChat, ModelID: "m", RenderedPrompt: "p"}) + var providerErr *Error + if !errors.As(err, &providerErr) || providerErr.Class != test.class { + t.Fatalf("error=%v", err) + } + }) + } +} + +func TestExtraBodyRejectsProtocolOverrides(t *testing.T) { + for _, value := range []string{`{"model":"override"}`, `{"messages":[]}`, `{"unknown":true}`} { + if _, err := NewOpenAI(&fakeHTTP{}, OpenAIConfig{BaseURL: "https://provider.test/v1", ExtraBody: jsonBytes(value)}); !errors.Is(err, ErrInvalidConfig) { + t.Fatalf("extra_body %s error=%v", value, err) + } + } +} + +func jsonBytes(value string) []byte { return []byte(value) } + +var _ HTTPClient = (*fakeHTTP)(nil) +var _ = http.MethodPost diff --git a/internal/core/provider/provider.go b/internal/core/provider/provider.go index d0505e5..b0c7c95 100644 --- a/internal/core/provider/provider.go +++ b/internal/core/provider/provider.go @@ -3,6 +3,7 @@ package provider import ( "context" "fmt" + "net/http" "git.ilapage.cn/OPC/chorus/internal/core/model" ) @@ -48,6 +49,7 @@ type Input struct { MIMEType string Role model.InputRole Position uint32 + Content []byte } type Request struct { @@ -68,3 +70,22 @@ type Output struct { type Client interface { Generate(ctx context.Context, request Request) ([]Output, error) } + +type HTTPRequest struct { + Method string + URL string + Header http.Header + Body []byte + MaxBytes int64 +} + +type HTTPResponse struct { + StatusCode int + ContentType string + Body []byte +} + +type HTTPClient interface { + Do(ctx context.Context, request HTTPRequest) (HTTPResponse, error) + Fetch(ctx context.Context, rawURL string, maxBytes int64) (HTTPResponse, error) +} diff --git a/internal/core/queue/mysql_repository.go b/internal/core/queue/mysql_repository.go index dcc4c25..cabff17 100644 --- a/internal/core/queue/mysql_repository.go +++ b/internal/core/queue/mysql_repository.go @@ -300,6 +300,17 @@ func (r *MySQLRepository) Fail(ctx context.Context, generationID uint64, leaseTo return result.RowsAffected == 1, nil } +func (r *MySQLRepository) Inputs(ctx context.Context, generationID uint64) ([]model.GenerationInput, error) { + if generationID == 0 { + return nil, ErrInvalidGeneration + } + var inputs []model.GenerationInput + if err := r.db.WithContext(ctx).Where("generation_id = ?", generationID).Order("position, id").Find(&inputs).Error; err != nil { + return nil, fmt.Errorf("load queue generation inputs: %w", err) + } + return inputs, nil +} + func newLeaseToken() (string, error) { bytes := make([]byte, 16) if _, err := rand.Read(bytes); err != nil { diff --git a/internal/platform/http/provider.go b/internal/platform/http/provider.go new file mode 100644 index 0000000..ca3e455 --- /dev/null +++ b/internal/platform/http/provider.go @@ -0,0 +1,27 @@ +package safehttp + +import ( + "bytes" + "context" + + "git.ilapage.cn/OPC/chorus/internal/core/provider" +) + +type ProviderClient struct{ client *Client } + +func NewProviderClient(client *Client) *ProviderClient { return &ProviderClient{client: client} } + +func (c *ProviderClient) Do(ctx context.Context, request provider.HTTPRequest) (provider.HTTPResponse, error) { + response, err := c.client.Do(ctx, Request{ + Method: request.Method, URL: request.URL, Header: request.Header, + Body: bytes.NewReader(request.Body), MaxBytes: request.MaxBytes, + }) + return provider.HTTPResponse{StatusCode: response.StatusCode, ContentType: response.ContentType, Body: response.Body}, err +} + +func (c *ProviderClient) Fetch(ctx context.Context, rawURL string, maxBytes int64) (provider.HTTPResponse, error) { + response, err := c.client.Fetch(ctx, rawURL, maxBytes) + return provider.HTTPResponse{StatusCode: response.StatusCode, ContentType: response.ContentType, Body: response.Body}, err +} + +var _ provider.HTTPClient = (*ProviderClient)(nil) diff --git a/internal/platform/http/safehttp.go b/internal/platform/http/safehttp.go index 6901976..7433141 100644 --- a/internal/platform/http/safehttp.go +++ b/internal/platform/http/safehttp.go @@ -18,6 +18,7 @@ var ( ErrInvalidTarget = errors.New("outbound target is invalid") ErrBlockedAddress = errors.New("outbound target address is blocked") ErrTooManyRedirects = errors.New("outbound redirect limit exceeded") + ErrCrossOriginAuth = errors.New("authenticated cross-origin redirect is blocked") ErrResponseTooLarge = errors.New("outbound response exceeds size limit") ErrRequestFailed = errors.New("safe outbound request failed") ) @@ -47,6 +48,14 @@ type Response struct { Body []byte } +type Request struct { + Method string + URL string + Header http.Header + Body io.Reader + MaxBytes int64 +} + func New(config Config) (*Client, error) { if config.Resolver == nil { config.Resolver = net.DefaultResolver @@ -87,26 +96,45 @@ func New(config Config) (*Client, error) { if len(via) > config.MaxRedirects { return ErrTooManyRedirects } + if len(via) > 0 && via[0].Header.Get("Authorization") != "" && !sameOrigin(via[0].URL, request.URL) { + return ErrCrossOriginAuth + } return v.validateURL(request.Context(), request.URL, true) } return &Client{httpClient: client, validator: v}, nil } +func sameOrigin(first, second *url.URL) bool { + if first == nil || second == nil { + return false + } + return strings.EqualFold(first.Scheme, second.Scheme) && strings.EqualFold(first.Host, second.Host) +} + func (c *Client) Fetch(ctx context.Context, rawURL string, maxBytes int64) (Response, error) { - if maxBytes <= 0 { + return c.Do(ctx, Request{Method: http.MethodGet, URL: rawURL, MaxBytes: maxBytes}) +} + +func (c *Client) Do(ctx context.Context, input Request) (Response, error) { + if input.MaxBytes <= 0 { return Response{}, fmt.Errorf("maximum response bytes must be positive") } - target, err := url.Parse(rawURL) + method := strings.ToUpper(strings.TrimSpace(input.Method)) + if method != http.MethodGet && method != http.MethodPost { + return Response{}, fmt.Errorf("safe HTTP method is not allowed") + } + target, err := url.Parse(input.URL) if err != nil { return Response{}, ErrInvalidTarget } if err := c.validator.validateURL(ctx, target, true); err != nil { return Response{}, err } - request, err := http.NewRequestWithContext(ctx, http.MethodGet, target.String(), nil) + request, err := http.NewRequestWithContext(ctx, method, target.String(), input.Body) if err != nil { return Response{}, ErrInvalidTarget } + request.Header = input.Header.Clone() response, err := c.httpClient.Do(request) if err != nil { var urlError *url.Error @@ -116,14 +144,14 @@ func (c *Client) Fetch(ctx context.Context, rawURL string, maxBytes int64) (Resp return Response{}, ErrRequestFailed } defer response.Body.Close() - if response.ContentLength > maxBytes { + if response.ContentLength > input.MaxBytes { return Response{}, ErrResponseTooLarge } - body, err := io.ReadAll(io.LimitReader(response.Body, maxBytes+1)) + body, err := io.ReadAll(io.LimitReader(response.Body, input.MaxBytes+1)) if err != nil { return Response{}, fmt.Errorf("read safe HTTP response: %w", err) } - if int64(len(body)) > maxBytes { + if int64(len(body)) > input.MaxBytes { return Response{}, ErrResponseTooLarge } return Response{StatusCode: response.StatusCode, ContentType: response.Header.Get("Content-Type"), Body: body}, nil diff --git a/internal/platform/http/safehttp_test.go b/internal/platform/http/safehttp_test.go index ef3bea3..1f94df0 100644 --- a/internal/platform/http/safehttp_test.go +++ b/internal/platform/http/safehttp_test.go @@ -106,6 +106,40 @@ func TestFetchPublicMockAndSizeLimit(t *testing.T) { } } +func TestDoPostUsesSafeTransportAndForwardsHeaders(t *testing.T) { + resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ + "provider.test": {ips("93.184.216.34"), ips("93.184.216.34")}, + }} + requestText := make(chan string, 1) + fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { + client, server := net.Pipe() + go func() { + defer server.Close() + buffer := make([]byte, 4096) + count, _ := server.Read(buffer) + requestText <- string(buffer[:count]) + _, _ = server.Write([]byte("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 11\r\nConnection: close\r\n\r\n{\"ok\":true}")) + }() + return client, nil + } + client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 0}) + if err != nil { + t.Fatal(err) + } + response, err := client.Do(context.Background(), Request{ + Method: http.MethodPost, URL: "http://provider.test/v1/chat/completions", + Header: http.Header{"Authorization": []string{"Bearer test-key"}, "Content-Type": []string{"application/json"}}, + Body: strings.NewReader(`{"model":"mock"}`), MaxBytes: 100, + }) + if err != nil || response.StatusCode != http.StatusOK { + t.Fatalf("Do() = %#v, %v", response, err) + } + received := <-requestText + if !strings.Contains(received, "POST /v1/chat/completions") || !strings.Contains(received, "Authorization: Bearer test-key") { + t.Fatalf("unexpected outbound request: %q", received) + } +} + func TestRedirectAndProxyBypassAreRejected(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ "public.test": {ips("93.184.216.34"), ips("93.184.216.34")}, @@ -136,6 +170,31 @@ func TestRedirectAndProxyBypassAreRejected(t *testing.T) { } } +func TestAuthenticatedCrossOriginRedirectIsRejected(t *testing.T) { + resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{ + "provider.test": {ips("93.184.216.34"), ips("93.184.216.34")}, + "other.test": {ips("93.184.216.35")}, + }} + fakeDial := func(_ context.Context, _, _ string) (net.Conn, error) { + client, server := net.Pipe() + go func() { + defer server.Close() + buffer := make([]byte, 4096) + _, _ = server.Read(buffer) + _, _ = server.Write([]byte("HTTP/1.1 302 Found\r\nLocation: http://other.test/result\r\nContent-Length: 0\r\nConnection: close\r\n\r\n")) + }() + return client, nil + } + client, err := New(Config{Resolver: resolver, DialContext: fakeDial, Timeout: time.Second, MaxRedirects: 2}) + if err != nil { + t.Fatal(err) + } + _, err = client.Do(context.Background(), Request{Method: http.MethodPost, URL: "http://provider.test/request", Header: http.Header{"Authorization": []string{"Bearer secret"}}, Body: strings.NewReader("{}"), MaxBytes: 100}) + if !errors.Is(err, ErrCrossOriginAuth) { + t.Fatalf("cross-origin redirect error = %v", err) + } +} + func TestURLRestrictionsAndRedaction(t *testing.T) { resolver := &fakeResolver{addresses: map[string][][]net.IPAddr{"public.test": {ips("93.184.216.34")}}} client, err := New(Config{Resolver: resolver, Timeout: time.Second, MaxRedirects: 0}) diff --git a/internal/platform/mockprovider/handler.go b/internal/platform/mockprovider/handler.go new file mode 100644 index 0000000..3447765 --- /dev/null +++ b/internal/platform/mockprovider/handler.go @@ -0,0 +1,133 @@ +package mockprovider + +import ( + "bytes" + "encoding/base64" + "encoding/json" + "image" + "image/color" + "image/png" + "net/http" + "strings" + "time" +) + +type Handler struct{ Delay time.Duration } + +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/edits": + h.image(response, request) + case "/v1/result.png": + response.Header().Set("Content-Type", "image/png") + response.Write(mockPNG()) + default: + http.NotFound(response, request) + } +} + +func (h Handler) chat(response http.ResponseWriter, request *http.Request) { + var body struct { + Messages []struct { + Content string `json:"content"` + } `json:"messages"` + } + if json.NewDecoder(http.MaxBytesReader(response, request.Body, 1<<20)).Decode(&body) != nil || len(body.Messages) == 0 { + writeError(response, 400, "bad_request") + return + } + prompt := body.Messages[0].Content + if h.scenario(response, request, prompt) { + return + } + response.Header().Set("Content-Type", "application/json") + _ = 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 + } + prompt := request.FormValue("prompt") + if h.scenario(response, request, prompt) { + return + } + files := request.MultipartForm.File["image[]"] + if len(files) == 0 { + writeError(response, 400, "bad_request") + return + } + for _, header := range files { + file, err := header.Open() + if err != nil { + writeError(response, 400, "bad_request") + return + } + _, err = png.Decode(file) + file.Close() + if err != nil { + writeError(response, 400, "bad_image") + return + } + } + data := map[string]any{"b64_json": base64.StdEncoding.EncodeToString(mockPNG())} + if strings.Contains(prompt, "mock:url") { + data = map[string]any{"url": "http://" + request.Host + "/v1/result.png"} + } + response.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(response).Encode(map[string]any{"data": []any{data}}) +} + +func (h Handler) scenario(response http.ResponseWriter, request *http.Request, prompt string) bool { + switch { + case strings.Contains(prompt, "mock:429"): + writeError(response, 429, "rate_limited") + case strings.Contains(prompt, "mock:500"): + writeError(response, 500, "server_error") + case strings.Contains(prompt, "mock:400"): + writeError(response, 400, "bad_request") + case strings.Contains(prompt, "mock:401"): + writeError(response, 401, "unauthorized") + case strings.Contains(prompt, "mock:policy"): + writeError(response, 400, "content_policy_violation") + case strings.Contains(prompt, "mock:timeout"): + delay := h.Delay + if delay <= 0 { + delay = time.Second + } + select { + case <-request.Context().Done(): + case <-time.After(delay): + writeError(response, 504, "timeout") + } + case strings.Contains(prompt, "mock:connection"): + if hijacker, ok := response.(http.Hijacker); ok { + connection, _, err := hijacker.Hijack() + if err == nil { + connection.Close() + } + } + default: + return false + } + return true +} +func writeError(response http.ResponseWriter, status int, code string) { + response.Header().Set("Content-Type", "application/json") + response.WriteHeader(status) + _ = json.NewEncoder(response).Encode(map[string]any{"error": map[string]string{"code": code, "message": "mock upstream failure"}}) +} +func mockPNG() []byte { + img := image.NewRGBA(image.Rect(0, 0, 2, 2)) + for y := 0; y < 2; y++ { + for x := 0; x < 2; x++ { + img.Set(x, y, color.RGBA{R: 40, G: 120, B: 200, A: 255}) + } + } + var output bytes.Buffer + _ = png.Encode(&output, img) + return output.Bytes() +} diff --git a/internal/platform/mockprovider/handler_test.go b/internal/platform/mockprovider/handler_test.go new file mode 100644 index 0000000..236ebf6 --- /dev/null +++ b/internal/platform/mockprovider/handler_test.go @@ -0,0 +1,43 @@ +package mockprovider + +import ( + "bytes" + "encoding/json" + "mime/multipart" + "net/http" + "net/http/httptest" + "testing" +) + +func TestChatSuccessAndErrorMatrix(t *testing.T) { + handler := Handler{} + for _, test := range []struct { + prompt string + status int + }{{"hello", 200}, {"mock:429", 429}, {"mock:500", 500}, {"mock:400", 400}, {"mock:401", 401}, {"mock:policy", 400}} { + t.Run(test.prompt, func(t *testing.T) { + body, _ := json.Marshal(map[string]any{"messages": []any{map[string]string{"content": test.prompt}}}) + request := httptest.NewRequest(http.MethodPost, "http://mock/v1/chat/completions", bytes.NewReader(body)) + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != test.status { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } + }) + } +} +func TestImageSuccess(t *testing.T) { + var body bytes.Buffer + writer := multipart.NewWriter(&body) + _ = writer.WriteField("prompt", "edit") + part, _ := writer.CreateFormFile("image[]", "input.png") + part.Write(mockPNG()) + writer.Close() + request := httptest.NewRequest(http.MethodPost, "http://mock/v1/images/edits", &body) + request.Header.Set("Content-Type", writer.FormDataContentType()) + response := httptest.NewRecorder() + Handler{}.ServeHTTP(response, request) + if response.Code != 200 || !bytes.Contains(response.Body.Bytes(), []byte("b64_json")) { + t.Fatalf("status=%d body=%s", response.Code, response.Body.String()) + } +} diff --git a/portal/worker/catalog.go b/portal/worker/catalog.go new file mode 100644 index 0000000..2434165 --- /dev/null +++ b/portal/worker/catalog.go @@ -0,0 +1,80 @@ +package worker + +import ( + "context" + "encoding/json" + "fmt" + "time" + + 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" + "gorm.io/gorm" +) + +type GORMCatalog struct { + db *gorm.DB + cipher corecrypto.KeyCipher +} + +func NewGORMCatalog(db *gorm.DB, cipher corecrypto.KeyCipher) (*GORMCatalog, error) { + if db == nil || cipher == nil { + return nil, ErrInvalidConfig + } + 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 { + return Selection{}, ErrNoProvider + } + selected := rows[0] + var apiKey string + switch selected.AuthType { + case "none": + case "bearer": + var envelope corecrypto.Envelope + if json.Unmarshal(selected.APIKeyEnc, &envelope) != nil { + return Selection{}, fmt.Errorf("decode provider credential") + } + plaintext, err := c.cipher.Decrypt(ctx, envelope) + if err != nil { + return Selection{}, fmt.Errorf("decrypt provider credential: %w", err) + } + apiKey = string(plaintext) + for i := range plaintext { + plaintext[i] = 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 +} + +type OpenAIFactory struct { + http provider.HTTPClient + maxResponseBytes int64 +} + +func NewOpenAIFactory(httpClient provider.HTTPClient, maxResponseBytes int64) (*OpenAIFactory, error) { + if httpClient == nil || maxResponseBytes <= 0 { + return nil, ErrInvalidConfig + } + 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}) +} diff --git a/portal/worker/local_storage.go b/portal/worker/local_storage.go new file mode 100644 index 0000000..2275aa9 --- /dev/null +++ b/portal/worker/local_storage.go @@ -0,0 +1,30 @@ +package worker + +import ( + "context" + "io" + + corestorage "git.ilapage.cn/OPC/chorus/internal/core/storage" + platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage" +) + +type LocalStorage struct{ storage *platformstorage.Local } + +func NewLocalStorage(storage *platformstorage.Local) (*LocalStorage, error) { + if storage == nil { + return nil, ErrInvalidConfig + } + return &LocalStorage{storage: storage}, nil +} +func (s *LocalStorage) Open(ctx context.Context, key string) (io.ReadCloser, corestorage.Object, error) { + return s.storage.Open(ctx, key) +} +func (s *LocalStorage) Delete(ctx context.Context, key string) error { + return s.storage.Delete(ctx, key) +} +func (s *LocalStorage) PutImage(ctx context.Context, request ImageRequest) (ImageObjects, error) { + objects, err := s.storage.PutImage(ctx, platformstorage.ImageRequest{Key: request.Key, ThumbnailKey: request.ThumbnailKey, OwnerID: request.OwnerID, GenerationID: request.GenerationID, ContentType: request.ContentType, Source: request.Source}) + return ImageObjects{Original: objects.Original, Thumbnail: objects.Thumbnail}, err +} + +var _ ImageStore = (*LocalStorage)(nil) diff --git a/portal/worker/mysql_integration_test.go b/portal/worker/mysql_integration_test.go new file mode 100644 index 0000000..724cfb0 --- /dev/null +++ b/portal/worker/mysql_integration_test.go @@ -0,0 +1,156 @@ +package worker + +import ( + "bytes" + "context" + "encoding/json" + "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" + 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" +) + +type fixedResolver struct{} + +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) { + 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{}) + if err != nil { + t.Fatal(err) + } + sqlDB, _ := db.DB() + t.Cleanup(func() { sqlDB.Close() }) + 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}}) + if err != nil { + t.Fatal(err) + } + providerHTTP := safehttp.NewProviderClient(httpClient) + factory, err := NewOpenAIFactory(providerHTTP, 2<<20) + if err != nil { + t.Fatal(err) + } + keyRing, err := platformcrypto.NewKeyRing("test", map[string][]byte{"test": bytes.Repeat([]byte{1}, 32)}) + if err != nil { + t.Fatal(err) + } + catalog, err := NewGORMCatalog(db, 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")) + 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) + if err != nil { + t.Fatal(err) + } + local, err := platformstorage.NewLocal(platformstorage.Config{Root: t.TempDir(), MaxObjectBytes: 2 << 20, MaxImagePixels: 10000, ThumbnailMaxSide: 32, AllowedImageMIME: map[string]bool{"image/png": true}}) + 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) + } + 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" + } + generation := model.Generation{UserID: userID, Kind: kind, IdempotencyKey: "worker-it-" + string(kind) + "-" + strconv.FormatInt(time.Now().UnixNano(), 10), UserPrompt: prompt, RenderedPrompt: prompt} + 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) + } + }) + } +} diff --git a/portal/worker/worker.go b/portal/worker/worker.go new file mode 100644 index 0000000..14f7300 --- /dev/null +++ b/portal/worker/worker.go @@ -0,0 +1,260 @@ +package worker + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "strings" + "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" + 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") +) + +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) + 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) +} + +type ClaimController interface{ StopClaims() } + +type Selection struct { + ProviderModelID uint64 + BaseURL string + APIKey string + ModelID string + APIType model.APIType + ExtraBody []byte + Timeout time.Duration +} + +type Catalog interface { + Select(ctx context.Context, kind model.GenerationKind) (Selection, error) +} +type ClientFactory interface { + New(selection Selection) (provider.Client, error) +} + +type ImageStore interface { + Open(ctx context.Context, key string) (io.ReadCloser, corestorage.Object, error) + PutImage(ctx context.Context, request ImageRequest) (ImageObjects, error) + Delete(ctx context.Context, key string) error +} + +type ImageRequest struct { + Key, ThumbnailKey string + OwnerID, GenerationID uint64 + ContentType string + Source io.Reader +} +type ImageObjects struct{ Original, Thumbnail corestorage.Object } + +type Config struct { + Owner string + LeaseDuration time.Duration + PollInterval time.Duration +} + +type Worker struct { + config Config + queue Queue + controller ClaimController + catalog Catalog + factory ClientFactory + 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 { + return nil, ErrInvalidConfig + } + return &Worker{config: config, queue: queueRepository, controller: controller, catalog: catalog, factory: factory, storage: storage}, nil +} + +func (w *Worker) Run(ctx context.Context) error { + stop := make(chan struct{}) + go func() { + select { + case <-ctx.Done(): + w.controller.StopClaims() + case <-stop: + } + }() + defer close(stop) + for { + if ctx.Err() != nil { + return nil + } + claim, err := w.queue.ClaimNext(ctx, w.config.Owner, w.config.LeaseDuration) + if errors.Is(err, queue.ErrClaimsStopped) || errors.Is(err, context.Canceled) { + return nil + } + if err != nil { + return fmt.Errorf("claim worker task: %w", err) + } + if claim == nil { + select { + case <-ctx.Done(): + return nil + case <-time.After(w.config.PollInterval): + continue + } + } + workCtx, cancel := context.WithTimeout(context.Background(), w.config.LeaseDuration) + if err := w.process(workCtx, claim); err != nil { + cancel() + return fmt.Errorf("process worker task: %w", err) + } + cancel() + } +} + +func (w *Worker) ProcessOne(ctx context.Context) (bool, error) { + claim, err := w.queue.ClaimNext(ctx, w.config.Owner, w.config.LeaseDuration) + if err != nil || claim == nil { + return false, err + } + return true, w.process(ctx, claim) +} + +func (w *Worker) process(ctx context.Context, claim *queue.Claim) error { + started := time.Now() + selection, err := w.catalog.Select(ctx, claim.Generation.Kind) + 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) + } + inputs, err := w.loadInputs(ctx, claim.Generation) + if err != nil { + return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) + } + requestCtx := ctx + var cancel context.CancelFunc + if selection.Timeout > 0 { + requestCtx, cancel = context.WithTimeout(ctx, selection.Timeout) + defer cancel() + } + generated, err := client.Generate(requestCtx, provider.Request{Kind: claim.Generation.Kind, APIType: selection.APIType, ModelID: selection.ModelID, RenderedPrompt: claim.Generation.RenderedPrompt, Inputs: inputs}) + 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, selection.ProviderModelID, started, provider.CodeUnknown, err) + } + outputs, keys, err := w.saveOutputs(ctx, claim.Generation, generated) + if err != nil { + w.deleteKeys(keys) + return w.fail(ctx, claim, selection.ProviderModelID, started, provider.CodeUnknown, err) + } + 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) + if err != nil { + return err + } + return nil + } + return nil +} + +func (w *Worker) loadInputs(ctx context.Context, generation model.Generation) ([]provider.Input, error) { + rows, err := w.queue.Inputs(ctx, generation.ID) + if err != nil { + return nil, err + } + inputs := make([]provider.Input, 0, len(rows)) + for _, row := range rows { + reader, object, openErr := w.storage.Open(ctx, row.StorageKey) + if openErr != nil { + return nil, openErr + } + content, readErr := io.ReadAll(reader) + closeErr := reader.Close() + if readErr != nil { + return nil, readErr + } + if closeErr != nil { + return nil, closeErr + } + if object.OwnerID != generation.UserID || object.GenerationID != generation.ID { + return nil, fmt.Errorf("input storage ownership mismatch") + } + inputs = append(inputs, provider.Input{StorageKey: row.StorageKey, MIMEType: row.MIMEType, Role: row.Role, Position: row.Position, Content: content}) + } + return inputs, nil +} + +func (w *Worker) saveOutputs(ctx context.Context, generation model.Generation, generated []provider.Output) ([]model.GenerationOutput, []string, error) { + if len(generated) == 0 { + return nil, nil, fmt.Errorf("provider returned no outputs") + } + rows := make([]model.GenerationOutput, 0, len(generated)) + keys := []string{} + for index, output := range generated { + if output.Text != "" { + text := output.Text + 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) + 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 { + return nil, keys, err + } + keys = append(keys, objects.Original.Key, objects.Thumbnail.Key) + mimeType := objects.Original.ContentType + size := uint64(objects.Original.Size) + originalKey, thumbnailKey := objects.Original.Key, objects.Thumbnail.Key + rows = append(rows, model.GenerationOutput{Kind: output.Kind, StorageKey: &originalKey, ThumbnailStorageKey: &thumbnailKey, MIMEType: &mimeType, SizeBytes: &size}) + } + 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} + 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 +} +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() + if elapsed < 1 { + return 1 + } + return elapsed +} diff --git a/portal/worker/worker_test.go b/portal/worker/worker_test.go new file mode 100644 index 0000000..81cd3c3 --- /dev/null +++ b/portal/worker/worker_test.go @@ -0,0 +1,241 @@ +package worker + +import ( + "bytes" + "context" + "image" + "image/color" + "image/png" + "sync" + "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" + 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 +} + +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) 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) Fail(_ context.Context, _ uint64, _ string, code, _ string, attempt model.Attempt) (bool, error) { + q.failedCode = code + q.attempt = attempt + return !q.stale, nil +} +func (q *fakeQueue) Inputs(context.Context, uint64) ([]model.GenerationInput, error) { + return q.inputs, nil +} + +type fakeController struct { + mutex sync.Mutex + stopped bool +} + +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 +} + +type fakeFactory struct{ client provider.Client } + +func (f fakeFactory) New(Selection) (provider.Client, error) { return f.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() + } + } + 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() + } + }) +} + +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) + _, err := w.ProcessOne(context.Background()) + if err != nil { + t.Fatal(err) + } + if _, _, err := store.Open(context.Background(), "outputs/1/10/1-1"); err == nil { + t.Fatal("stale output was retained") + } +} + +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) + _, 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) + } +} + +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 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 { + t.Helper() + w, err := New(Config{Owner: "test", LeaseDuration: time.Second, PollInterval: time.Millisecond}, q, c, catalog, factory, 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++ { + for x := 0; x < 4; x++ { + img.Set(x, y, color.RGBA{R: 200, G: 80, B: 40, A: 255}) + } + } + var output bytes.Buffer + if err := png.Encode(&output, img); err != nil { + panic(err) + } + return output.Bytes() +}