feat: add provider protocols and embedded worker (#10)

This commit is contained in:
ila
2026-08-21 00:18:19 +08:00
parent d61d73d96f
commit d61c80265b
15 changed files with 1490 additions and 6 deletions
+22
View File
@@ -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)
}
}
+264
View File
@@ -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)
+109
View File
@@ -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
+21
View File
@@ -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)
}
+11
View File
@@ -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 {
+27
View File
@@ -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)
+34 -6
View File
@@ -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
+59
View File
@@ -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})
+133
View File
@@ -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()
}
@@ -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())
}
}
+80
View File
@@ -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})
}
+30
View File
@@ -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)
+156
View File
@@ -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)
}
})
}
}
+260
View File
@@ -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
}
+241
View File
@@ -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()
}