feat: add provider protocols and embedded worker (#10)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
Reference in New Issue
Block a user