259 lines
12 KiB
Go
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)
|
|
}
|
|
}
|