Files
chorus/portal/worker/mysql_integration_test.go
T

259 lines
12 KiB
Go

package worker
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net"
"net/http/httptest"
"os"
"testing"
"time"
"git.ilapage.cn/OPC/chorus/internal/core/model"
"git.ilapage.cn/OPC/chorus/internal/core/queue"
"git.ilapage.cn/OPC/chorus/internal/core/router"
platformcrypto "git.ilapage.cn/OPC/chorus/internal/platform/crypto"
safehttp "git.ilapage.cn/OPC/chorus/internal/platform/http"
"git.ilapage.cn/OPC/chorus/internal/platform/mockprovider"
platformstorage "git.ilapage.cn/OPC/chorus/internal/platform/storage"
"gorm.io/driver/mysql"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
type fixedResolver struct{}
func (fixedResolver) LookupIPAddr(context.Context, string) ([]net.IPAddr, error) {
return []net.IPAddr{{IP: net.ParseIP("93.184.216.34")}}, nil
}
func TestMySQLWorkerUsesSnapshotAndMockUpstream(t *testing.T) {
dsn := os.Getenv("CHORUS_TEST_DSN")
if dsn == "" {
t.Skip("CHORUS_TEST_DSN is not set")
}
db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatal(err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
defer sqlDB.Close()
tx := db.Begin()
if tx.Error != nil {
t.Fatal(tx.Error)
}
defer tx.Rollback()
server := httptest.NewServer(mockprovider.Handler{})
defer server.Close()
httpClient, err := safehttp.New(safehttp.Config{
Resolver: fixedResolver{}, Timeout: 2 * time.Second, MaxRedirects: 1, AllowedPorts: []uint16{80},
DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) {
return (&net.Dialer{}).DialContext(ctx, network, server.Listener.Addr().String())
},
})
if err != nil {
t.Fatal(err)
}
factory, err := NewOpenAIFactory(safehttp.NewProviderClient(httpClient), 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(tx, keyRing)
if err != nil {
t.Fatal(err)
}
runtime, err := NewGORMRuntime(tx)
if err != nil {
t.Fatal(err)
}
queueRepository, err := queue.NewMySQLRepository(tx)
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, err := NewLocalStorage(local)
if err != nil {
t.Fatal(err)
}
suffix := time.Now().UnixNano()
user := model.User{Email: fmt.Sprintf("worker-%d@chorus.invalid", suffix), PasswordHash: "synthetic", DisplayName: "Worker", Status: "active"}
providerRow := model.Provider{Slug: fmt.Sprintf("worker-%d", suffix), Name: "Worker Mock", BaseURL: "http://provider.test/v1", AuthType: "none", Enabled: true}
modelRow := model.ProviderModel{ProviderID: 0, Name: "Worker Chat", ModelID: "mock-chat", APIType: model.APIChat, Kind: model.KindText, ExtraBody: json.RawMessage("{}"), TimeoutMS: 1000, Weight: 1, Enabled: true}
template := model.PromptTemplate{TemplateKey: fmt.Sprintf("worker-%d", suffix), Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Worker", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true}
if err := tx.Create(&user).Error; err != nil {
t.Fatal(err)
}
if err := tx.Create(&providerRow).Error; err != nil {
t.Fatal(err)
}
modelRow.ProviderID = providerRow.ID
if err := tx.Create(&modelRow).Error; err != nil {
t.Fatal(err)
}
if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id, capability) VALUES (?, ?)", modelRow.ID, model.CapabilityText).Error; err != nil {
t.Fatal(err)
}
if err := tx.Create(&template).Error; err != nil {
t.Fatal(err)
}
if err := tx.Exec("INSERT INTO route_pools (slug, name, capability, prompt_template_id, max_failover, version, enabled) VALUES (?, ?, ?, ?, 0, 1, TRUE)", fmt.Sprintf("worker-%d", suffix), "Worker", model.CapabilityText, template.ID).Error; err != nil {
t.Fatal(err)
}
var poolID uint64
if err := tx.Raw("SELECT id FROM route_pools WHERE slug = ?", fmt.Sprintf("worker-%d", suffix)).Scan(&poolID).Error; err != nil || poolID == 0 {
t.Fatalf("route pool id=%d error=%v", poolID, err)
}
if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id, provider_model_id, weight, failure_threshold, open_seconds, half_open_max, enabled) VALUES (?, ?, 1, 2, 60, 1, TRUE)", poolID, modelRow.ID).Error; err != nil {
t.Fatal(err)
}
var memberID uint64
if err := tx.Raw("SELECT id FROM route_pool_members WHERE route_pool_id = ?", poolID).Scan(&memberID).Error; err != nil || memberID == 0 {
t.Fatalf("route member id=%d error=%v", memberID, err)
}
snapshot := router.RouteSnapshot{Capability: model.CapabilityText, RoutePoolID: poolID, RoutePoolVersion: 1, PromptTemplateID: template.ID, PromptTemplateKey: template.TemplateKey, PromptTemplateVersion: 1, Members: []router.MemberSnapshot{{RoutePoolMemberID: memberID, ProviderModelID: modelRow.ID, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}}}
generation := model.Generation{UserID: user.ID, Kind: model.KindText, IdempotencyKey: fmt.Sprintf("worker-%d", suffix), UserPrompt: "integration", RenderedPrompt: "integration", CreatedAt: time.Unix(1, 0)}
if err := router.ApplySnapshot(&generation, snapshot); err != nil {
t.Fatal(err)
}
created, err := queueRepository.CreateIdempotent(context.Background(), &generation, nil)
if err != nil || !created {
t.Fatalf("create=%v error=%v", created, err)
}
controller := queue.NewController(queueRepository)
w, err := New(Config{Owner: "integration", LeaseDuration: 5 * time.Second, PollInterval: time.Millisecond, Random: firstRandom{}}, queueRepository, controller, catalog, factory, runtime, store)
if err != nil {
t.Fatal(err)
}
worked, err := w.ProcessOne(context.Background())
if err != nil || !worked {
t.Fatalf("worked=%v error=%v", worked, err)
}
var stored model.Generation
if err := tx.First(&stored, generation.ID).Error; err != nil {
t.Fatal(err)
}
if stored.Status != model.StatusSucceeded || stored.ProviderAttemptCount != 1 || !bytes.Contains(stored.Attempts, []byte(`"route_member_id"`)) || !bytes.Contains(stored.Attempts, []byte(`"finished_at"`)) {
t.Fatalf("stored generation=%#v attempts=%s", stored, stored.Attempts)
}
var outputCount int64
if err := tx.Model(&model.GenerationOutput{}).Where("generation_id = ?", generation.ID).Count(&outputCount).Error; err != nil || outputCount != 1 {
t.Fatalf("outputs=%d error=%v", outputCount, err)
}
}
func TestMySQLRuntimeOpensAndLimitsHalfOpenProbe(t *testing.T) {
dsn := os.Getenv("CHORUS_TEST_DSN")
if dsn == "" {
t.Skip("CHORUS_TEST_DSN is not set")
}
db, err := gorm.Open(mysql.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
if err != nil {
t.Fatal(err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatal(err)
}
defer sqlDB.Close()
tx := db.Begin()
if tx.Error != nil {
t.Fatal(tx.Error)
}
defer tx.Rollback()
suffix := time.Now().UnixNano()
providerRow := model.Provider{Slug: fmt.Sprintf("runtime-%d", suffix), Name: "Runtime", BaseURL: "https://provider.invalid/v1", AuthType: "none", Enabled: true}
if err := tx.Create(&providerRow).Error; err != nil {
t.Fatal(err)
}
modelRow := model.ProviderModel{ProviderID: providerRow.ID, Name: "Runtime", ModelID: "runtime", APIType: model.APIChat, Kind: model.KindText, ExtraBody: json.RawMessage("{}"), TimeoutMS: 1000, Weight: 1, Enabled: true}
if err := tx.Create(&modelRow).Error; err != nil {
t.Fatal(err)
}
if err := tx.Exec("INSERT INTO provider_model_capabilities (provider_model_id, capability) VALUES (?, ?)", modelRow.ID, model.CapabilityText).Error; err != nil {
t.Fatal(err)
}
template := model.PromptTemplate{TemplateKey: fmt.Sprintf("runtime-%d", suffix), Kind: model.KindText, APIType: model.APIChat, Capability: model.CapabilityText, Name: "Runtime", Version: 1, TemplateText: "{{.UserPrompt}}", DefaultRoleRule: "", Enabled: true}
if err := tx.Create(&template).Error; err != nil {
t.Fatal(err)
}
poolSlug := fmt.Sprintf("runtime-%d", suffix)
if err := tx.Exec("INSERT INTO route_pools (slug, name, capability, prompt_template_id, max_failover, version, enabled) VALUES (?, 'Runtime', ?, ?, 0, 1, TRUE)", poolSlug, model.CapabilityText, template.ID).Error; err != nil {
t.Fatal(err)
}
var poolID uint64
if err := tx.Raw("SELECT id FROM route_pools WHERE slug = ?", poolSlug).Scan(&poolID).Error; err != nil || poolID == 0 {
t.Fatalf("pool id=%d error=%v", poolID, err)
}
if err := tx.Exec("INSERT INTO route_pool_members (route_pool_id, provider_model_id, weight, failure_threshold, open_seconds, half_open_max, enabled) VALUES (?, ?, 1, 2, 60, 1, TRUE)", poolID, modelRow.ID).Error; err != nil {
t.Fatal(err)
}
var memberID uint64
if err := tx.Raw("SELECT id FROM route_pool_members WHERE route_pool_id = ?", poolID).Scan(&memberID).Error; err != nil || memberID == 0 {
t.Fatalf("member id=%d error=%v", memberID, err)
}
runtime, err := NewGORMRuntime(tx)
if err != nil {
t.Fatal(err)
}
members := []router.MemberSnapshot{{RoutePoolMemberID: memberID, ProviderModelID: modelRow.ID, Weight: 1, FailureThreshold: 2, OpenSeconds: 60, HalfOpenMax: 1}}
state := func() router.MemberState {
states, stateErr := runtime.MemberStates(context.Background(), model.CapabilityText, members)
if stateErr != nil || len(states) != 1 {
t.Fatalf("states=%#v error=%v", states, stateErr)
}
return states[0]
}
if current := state(); current.CircuitState != router.CircuitClosed || !current.Eligible() {
t.Fatalf("initial state=%#v", current)
}
for attempt := 0; attempt < 2; attempt++ {
reservation, reserveErr := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "runtime-test", LeaseDuration: time.Second})
if reserveErr != nil || reservation.State != router.CircuitClosed {
t.Fatalf("closed reservation=%#v error=%v", reservation, reserveErr)
}
owned, recordErr := runtime.Record(context.Background(), reservation, router.CircuitObservation{Retryable: true, ErrorCode: "upstream_timeout", ErrorMessage: "upstream request failed"})
if recordErr != nil || !owned {
t.Fatalf("record retryable owned=%v error=%v", owned, recordErr)
}
}
if current := state(); current.CircuitState != router.CircuitOpen || current.Eligible() {
t.Fatalf("open state=%#v", current)
}
if err := tx.Exec("UPDATE route_member_runtime SET open_until = ? WHERE route_pool_member_id = ?", time.Now().Add(-time.Second), memberID).Error; err != nil {
t.Fatal(err)
}
if current := state(); current.CircuitState != router.CircuitHalfOpen || current.HalfOpenInFlight != 0 || !current.Eligible() {
t.Fatalf("ready half-open state=%#v", current)
}
probe, err := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "runtime-test", LeaseDuration: time.Second})
if err != nil || probe.State != router.CircuitHalfOpen || probe.Token == "" {
t.Fatalf("probe=%#v error=%v", probe, err)
}
if _, err := runtime.Reserve(context.Background(), router.ReservationRequest{RoutePoolMemberID: memberID, Capability: model.CapabilityText, Owner: "second-worker", LeaseDuration: time.Second}); !errors.Is(err, router.ErrMemberUnavailable) {
t.Fatalf("second probe error=%v", err)
}
if owned, err := runtime.Record(context.Background(), probe, router.CircuitObservation{Succeeded: true}); err != nil || !owned {
t.Fatalf("successful probe owned=%v error=%v", owned, err)
}
if current := state(); current.CircuitState != router.CircuitClosed || current.HalfOpenInFlight != 0 || !current.Eligible() {
t.Fatalf("closed after probe state=%#v", current)
}
}