feat(worker): 实现集群限流与自适应负载

保留平台模型 RPM、TPM 和并发策略语义,增加 PostgreSQL 集群级租约、饱和候选重选和多平台自动负载,避免突发任务固定等待首个平台。\n\n新增 Worker 实时负载采样、自适应 active/heavy 容量、心跳与管理端指标,并扩展本地 acceptance runner,覆盖三 Worker、同模型三平台 2/4/6 并发和 48 个带图视频突发任务。\n\n验证:go test ./...、go vet ./...、PostgreSQL 跨 Store 集成测试、gofmt、bash -n、ShellCheck 及本地集群 provider-burst 验收通过;48/48 成功,无越限、重复提交、重复计费、重复回调或终态资源泄漏。
This commit is contained in:
2026-08-03 00:13:46 +08:00
parent 9a01fd4657
commit c28bf74230
52 changed files with 3700 additions and 272 deletions
+184 -35
View File
@@ -38,8 +38,10 @@ const (
)
const (
geminiRequestPrefix = `{"contents":[{"role":"user","parts":[{"text":"将参考图编辑为蓝色赛博朋克风格,保留主体结构"},{"inlineData":{"mimeType":"image/png","data":"`
geminiRequestSuffix = `"}}]}],"generationConfig":{"responseModalities":["IMAGE"]}}`
geminiRequestPrefix = `{"contents":[{"role":"user","parts":[{"text":"将多张参考图融合编辑为蓝色赛博朋克风格,保留主体结构"}`
geminiRequestImageStart = `,{"inlineData":{"mimeType":"image/png","data":"`
geminiRequestImageEnd = `"}}`
geminiRequestSuffix = `]}],"generationConfig":{"responseModalities":["IMAGE"]}}`
)
var reportURLPattern = regexp.MustCompile(`(?i)(?:https?|postgres(?:ql)?):\/\/[^\s"'<>]+`)
@@ -64,6 +66,8 @@ type options struct {
shardIndex int
shardCount int
executionID string
requestCount int
externalClient *http.Client
}
type report struct {
@@ -91,6 +95,7 @@ type phaseReport struct {
P99MS float64 `json:"p99Ms"`
DecodedOutputBytes int64 `json:"decodedOutputBytes,omitempty"`
OutputSHA256 string `json:"outputSha256,omitempty"`
InputImagesPerRequest int `json:"inputImagesPerRequest,omitempty"`
UniqueTaskIDs int `json:"uniqueTaskIds,omitempty"`
UniqueImageCombos int `json:"uniqueImageCombinations,omitempty"`
ForcedConversionTasks int `json:"forcedConversionTasks,omitempty"`
@@ -148,7 +153,7 @@ func main() {
ctx, cancel := context.WithTimeout(signalContext, opts.timeout)
defer cancel()
result := report{
SchemaVersion: "acceptance-load-report/v1",
SchemaVersion: "acceptance-load-report/v2",
RunID: opts.runID, Profile: opts.profile, StartedAt: time.Now().UTC(), SecretSafe: true,
}
phaseResults, runErr := run(ctx, opts)
@@ -205,6 +210,7 @@ func parseOptions() (options, error) {
var shardIndex int
var shardCount int
var executionID string
var requestCount int
flag.StringVar(&profile, "profile", env("AI_GATEWAY_ACCEPTANCE_PROFILE", "simulated-all"), "simulated-all, Gemini profile, video profile, or real-canary")
flag.StringVar(&reportPath, "report", env("AI_GATEWAY_ACCEPTANCE_REPORT", ""), "secret-safe JSON report path")
flag.DurationVar(&timeout, "timeout", 45*time.Minute, "overall timeout")
@@ -214,6 +220,7 @@ func parseOptions() (options, error) {
flag.IntVar(&shardIndex, "shard-index", 0, "zero-based distributed load shard index")
flag.IntVar(&shardCount, "shard-count", 1, "distributed load shard count")
flag.StringVar(&executionID, "execution-id", "", "unique workload execution identifier")
flag.IntVar(&requestCount, "requests", 0, "override the fixed profile request count")
flag.Parse()
opts := options{
gateways: splitCSV(os.Getenv("AI_GATEWAY_ACCEPTANCE_GATEWAYS")),
@@ -235,6 +242,7 @@ func parseOptions() (options, error) {
shardIndex: shardIndex,
shardCount: shardCount,
executionID: strings.ToLower(strings.TrimSpace(executionID)),
requestCount: requestCount,
}
if opts.executionID == "" {
opts.executionID = opts.profile
@@ -245,8 +253,8 @@ func parseOptions() (options, error) {
if opts.shardCount < 1 || opts.shardCount > 32 || opts.shardIndex < 0 || opts.shardIndex >= opts.shardCount {
return options{}, fmt.Errorf("invalid load shard %d/%d", opts.shardIndex, opts.shardCount)
}
if len(opts.gateways) < 1 || len(opts.gateways) > 2 || (opts.shardCount == 1 && len(opts.gateways) != 2) {
return options{}, errors.New("AI_GATEWAY_ACCEPTANCE_GATEWAYS must contain two API base URLs, or one URL for a distributed shard")
if len(opts.gateways) < 1 || len(opts.gateways) > 2 {
return options{}, errors.New("AI_GATEWAY_ACCEPTANCE_GATEWAYS must contain one or two API base URLs")
}
for index := range opts.gateways {
opts.gateways[index] = strings.TrimRight(opts.gateways[index], "/")
@@ -267,8 +275,11 @@ func parseOptions() (options, error) {
if opts.timeout <= 0 {
return options{}, errors.New("timeout must be positive")
}
if opts.requestCount < 0 || opts.requestCount > 10000 {
return options{}, errors.New("requests must be between 0 and 10000")
}
switch opts.profile {
case "simulated-smoke", "simulated-all", "gemini-baseline", "gemini-large", "gemini-peak", "video-throughput", "video-recovery",
case "simulated-smoke", "simulated-all", "gemini-baseline", "gemini-multi-image", "gemini-large", "gemini-peak", "video-throughput", "video-recovery",
"mixed-soak", "mixed-overload":
if opts.emulatorURL == "" {
return options{}, errors.New("AI_GATEWAY_ACCEPTANCE_EMULATOR_URL is required for simulated profiles")
@@ -322,28 +333,44 @@ func run(ctx context.Context, opts options) ([]phaseResult, error) {
TLSClientConfig: tlsConfig,
},
}
opts.externalClient = &http.Client{
Timeout: opts.timeout,
Transport: &http.Transport{
MaxIdleConns: 512,
MaxIdleConnsPerHost: 256,
ForceAttemptHTTP2: true,
},
}
results := make([]phaseResult, 0, 5)
runPhase := func(name string) error {
var result phaseResult
if profile, ok := acceptanceworkload.GeminiProfileByName(name); ok {
requests := profile.Requests
if opts.requestCount > 0 {
requests = opts.requestCount
}
result = runGemini(
ctx, client, opts, name, opts.shardRequestCount(profile.Requests),
profile.InputBytes, profile.OutputBytes, false, opts.shardIndex, opts.shardCount,
ctx, client, opts, name, opts.shardRequestCount(requests),
profile.InputImages, profile.InputBytes, profile.OutputBytes, false, opts.shardIndex, opts.shardCount,
)
} else {
switch name {
case "smoke-gemini":
result = runGemini(ctx, client, opts, name, opts.shardRequestCount(4), 256<<10, 256<<10, false, opts.shardIndex, opts.shardCount)
result = runGemini(ctx, client, opts, name, opts.shardRequestCount(4), 1, 256<<10, 256<<10, false, opts.shardIndex, opts.shardCount)
case "smoke-video":
result = runVideo(ctx, client, opts, name, opts.shardRequestCount(10), false, opts.emulatorFixtureURLs(), false, opts.shardIndex, opts.shardCount)
case "video-throughput":
result = runVideo(ctx, client, opts, name, opts.shardRequestCount(1200), false, opts.emulatorFixtureURLs(), false, opts.shardIndex, opts.shardCount)
requests := 1200
if opts.requestCount > 0 {
requests = opts.requestCount
}
result = runVideo(ctx, client, opts, name, opts.shardRequestCount(requests), false, opts.emulatorFixtureURLs(), false, opts.shardIndex, opts.shardCount)
case "video-recovery":
result = runVideo(ctx, client, opts, name, opts.shardRequestCount(96), true, opts.emulatorFixtureURLs(), false, opts.shardIndex, opts.shardCount)
case "mixed-soak", "mixed-overload":
result = runMixed(ctx, client, opts, name == "mixed-overload")
case "real-gemini-canary":
result = runGemini(ctx, client, opts, name, 1, 256<<10, 0, true, 0, 1)
result = runGemini(ctx, client, opts, name, 1, 1, 256<<10, 0, true, 0, 1)
case "real-video-canary":
result = runVideo(ctx, client, opts, name, 1, false, opts.realImageURLs, true, 0, 1)
default:
@@ -375,6 +402,7 @@ func runGemini(
opts options,
name string,
requestCount int,
inputImages int,
inputBytes int,
expectedOutputBytes int,
realUpstream bool,
@@ -412,8 +440,14 @@ func runGemini(
defer func() { <-slots }()
requestStarted := time.Now()
endpoint := opts.gateways[logicalIndex%len(opts.gateways)] + "/v1beta/models/" + url.PathEscape(opts.geminiModel) + ":generateContent"
input := paddedPNGVariant(inputBytes, geminiInputVariant(opts.runID+"/"+opts.executionID, logicalIndex))
req, err := newGeminiRequest(ctx, endpoint, input)
var req *http.Request
var err error
if name == acceptanceworkload.GeminiMultiImage.Name {
req, err = newGeminiFileDataRequest(ctx, endpoint, geminiFixtureURLs(opts, inputImages, logicalIndex))
} else {
inputs := geminiInputs(inputImages, inputBytes, opts.runID+"/"+opts.executionID, logicalIndex)
req, err = newGeminiRequest(ctx, endpoint, inputs)
}
if err == nil {
opts.setHeaders(req, logicalIndex, realUpstream)
req.Header.Set("Content-Type", "application/json")
@@ -462,6 +496,7 @@ func runGemini(
result := phaseReport{
Name: name, Requests: requestCount, Completed: completed, Failed: failures,
DurationMS: elapsed.Milliseconds(), DecodedOutputBytes: decodedBytes, OutputSHA256: expectedHash,
InputImagesPerRequest: inputImages,
}
if firstErr == nil && name == "gemini-baseline" && elapsed > 8*time.Minute {
firstErr = fmt.Errorf("Gemini baseline exceeded 8 minutes: %s", elapsed)
@@ -472,7 +507,7 @@ func runGemini(
return phaseResult{report: result, latencies: latencies, err: firstErr}
}
func streamGeminiRequestBody(input []byte) io.ReadCloser {
func streamGeminiRequestBody(inputs [][]byte) io.ReadCloser {
reader, writer := io.Pipe()
go func() {
closeWithError := func(err error) {
@@ -482,14 +517,24 @@ func streamGeminiRequestBody(input []byte) io.ReadCloser {
closeWithError(err)
return
}
encoder := base64.NewEncoder(base64.StdEncoding, writer)
if _, err := encoder.Write(input); err != nil {
closeWithError(err)
return
}
if err := encoder.Close(); err != nil {
closeWithError(err)
return
for _, input := range inputs {
if _, err := io.WriteString(writer, geminiRequestImageStart); err != nil {
closeWithError(err)
return
}
encoder := base64.NewEncoder(base64.StdEncoding, writer)
if _, err := encoder.Write(input); err != nil {
closeWithError(err)
return
}
if err := encoder.Close(); err != nil {
closeWithError(err)
return
}
if _, err := io.WriteString(writer, geminiRequestImageEnd); err != nil {
closeWithError(err)
return
}
}
if _, err := io.WriteString(writer, geminiRequestSuffix); err != nil {
closeWithError(err)
@@ -500,9 +545,9 @@ func streamGeminiRequestBody(input []byte) io.ReadCloser {
return reader
}
func newGeminiRequest(ctx context.Context, endpoint string, input []byte) (*http.Request, error) {
func newGeminiRequest(ctx context.Context, endpoint string, inputs [][]byte) (*http.Request, error) {
getBody := func() (io.ReadCloser, error) {
return streamGeminiRequestBody(input), nil
return streamGeminiRequestBody(inputs), nil
}
body, _ := getBody()
request, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, body)
@@ -511,10 +556,57 @@ func newGeminiRequest(ctx context.Context, endpoint string, input []byte) (*http
return nil, err
}
request.GetBody = getBody
request.ContentLength = int64(len(geminiRequestPrefix) + base64.StdEncoding.EncodedLen(len(input)) + len(geminiRequestSuffix))
contentLength := len(geminiRequestPrefix) + len(geminiRequestSuffix)
for _, input := range inputs {
contentLength += len(geminiRequestImageStart) + base64.StdEncoding.EncodedLen(len(input)) + len(geminiRequestImageEnd)
}
request.ContentLength = int64(contentLength)
return request, nil
}
func newGeminiFileDataRequest(ctx context.Context, endpoint string, imageURLs []string) (*http.Request, error) {
parts := make([]any, 0, len(imageURLs)+1)
parts = append(parts, map[string]any{"text": "将多张参考图融合编辑为蓝色赛博朋克风格,保留主体结构"})
for _, imageURL := range imageURLs {
parts = append(parts, map[string]any{"fileData": map[string]any{
"mimeType": geminiFixtureMIMEType(imageURL),
"fileUri": imageURL,
}})
}
body, err := json.Marshal(map[string]any{
"contents": []any{map[string]any{"role": "user", "parts": parts}},
"generationConfig": map[string]any{"responseModalities": []string{"IMAGE"}},
})
if err != nil {
return nil, err
}
return http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body))
}
func geminiFixtureMIMEType(imageURL string) string {
lower := strings.ToLower(imageURL)
switch {
case strings.Contains(lower, ".webp"):
return "image/webp"
case strings.Contains(lower, ".jpg"), strings.Contains(lower, ".jpeg"):
return "image/jpeg"
default:
return "image/png"
}
}
func geminiFixtureURLs(opts options, count int, logicalIndex int) []string {
fixtures := opts.emulatorFixtureURLs()
if count > len(fixtures) {
count = len(fixtures)
}
urls := make([]string, 0, count)
for imageIndex := 0; imageIndex < count; imageIndex++ {
urls = append(urls, fixtures[(logicalIndex*count+imageIndex)%len(fixtures)])
}
return urls
}
func runVideo(
ctx context.Context,
client *http.Client,
@@ -629,7 +721,7 @@ func runVideo(
}
wg.Wait()
submissionDuration := time.Since(submitStartedAt)
if firstErr == nil && name == "video-throughput" && submissionDuration > 10*time.Second {
if firstErr == nil && name == "video-throughput" && requestCount == 1200 && submissionDuration > 10*time.Second {
firstErr = fmt.Errorf("1200 video submissions exceeded 10 seconds: %s", submissionDuration)
}
uniqueTaskIDs := map[string]struct{}{}
@@ -672,18 +764,25 @@ func runVideo(
}()
}
wg.Wait()
comboSet := map[string]struct{}{}
for _, combo := range combinations {
comboSet[strings.Join(combo, "\x00")] = struct{}{}
}
if firstErr == nil && !realUpstream && len(comboSet) != 128 {
firstErr = fmt.Errorf("video image combinations=%d, want 128", len(comboSet))
uniqueImageCombinations := requestedVideoCombinationCount(
combinations,
requestCount,
requestOffset,
requestStride,
)
expectedCombinations := min(requestCount, 128)
if firstErr == nil && !realUpstream && uniqueImageCombinations < expectedCombinations {
firstErr = fmt.Errorf(
"video image combinations=%d, want at least %d",
uniqueImageCombinations,
expectedCombinations,
)
}
elapsed := time.Since(startedAt)
result := phaseReport{
Name: name, Requests: requestCount, Completed: completed, Failed: failures,
DurationMS: elapsed.Milliseconds(), SubmissionDurationMS: submissionDuration.Milliseconds(),
UniqueTaskIDs: len(uniqueTaskIDs), UniqueImageCombos: len(comboSet),
UniqueTaskIDs: len(uniqueTaskIDs), UniqueImageCombos: uniqueImageCombinations,
ForcedConversionTasks: forcedConversions,
}
if firstErr == nil && completed != requestCount {
@@ -755,7 +854,7 @@ func runMixed(
defer func() { <-slots }()
var result phaseResult
if isImage {
result = runGemini(runCtx, client, opts, "mixed-image", 1, 256<<10, 256<<10, false, requestIndex, 1)
result = runGemini(runCtx, client, opts, "mixed-image", 1, 1, 256<<10, 256<<10, false, requestIndex, 1)
} else {
result = runVideo(
runCtx,
@@ -931,7 +1030,14 @@ func validateVideoAsset(ctx context.Context, client *http.Client, opts options,
if gatewayRequest {
opts.setHeaders(request, index, false)
}
response, err := client.Do(request)
mediaClient := client
if !gatewayRequest && opts.gatewayTLSName != "" {
mediaClient = opts.externalClient
if mediaClient == nil {
mediaClient = &http.Client{Timeout: client.Timeout}
}
}
response, err := mediaClient.Do(request)
if err != nil {
return fmt.Errorf("download final video: %w", err)
}
@@ -1056,6 +1162,31 @@ func videoCombinations(images []string, count int) [][]string {
return out
}
func requestedVideoCombinationCount(
combinations [][]string,
requestCount int,
requestOffset int,
requestStride int,
) int {
if len(combinations) == 0 || requestCount <= 0 {
return 0
}
if requestStride < 1 {
requestStride = 1
}
seen := map[string]struct{}{}
for index := 0; index < requestCount; index++ {
logicalIndex := requestOffset + index*requestStride
imageCount := acceptanceworkload.VideoImageCount(logicalIndex)
combo := combinations[logicalIndex%len(combinations)]
if len(combo) > imageCount {
combo = combo[:imageCount]
}
seen[strings.Join(combo, "\x00")] = struct{}{}
}
return len(seen)
}
func streamGeminiImageHash(reader io.Reader) (int64, string, error) {
buffered := bufio.NewReaderSize(reader, 64<<10)
if err := scanUntil(buffered, []byte(`"inlineData"`), 256<<20); err != nil {
@@ -1176,6 +1307,24 @@ func paddedPNG(size int) []byte {
return paddedPNGVariant(size, 0)
}
func geminiInputs(count int, totalBytes int, executionID string, logicalIndex int) [][]byte {
if count < 1 {
count = 1
}
inputs := make([][]byte, 0, count)
baseSize := totalBytes / count
remainder := totalBytes % count
for imageIndex := 0; imageIndex < count; imageIndex++ {
size := baseSize
if imageIndex < remainder {
size++
}
variant := geminiInputVariant(executionID, logicalIndex*count+imageIndex)
inputs = append(inputs, paddedPNGVariant(size, variant))
}
return inputs
}
func geminiInputVariant(runID string, logicalIndex int) int {
digest := sha256.Sum256([]byte(fmt.Sprintf("%s:%d", strings.TrimSpace(runID), logicalIndex)))
return int(binary.BigEndian.Uint64(digest[:8]) & uint64(^uint(0)>>1))
+74 -17
View File
@@ -17,6 +17,12 @@ import (
"testing"
)
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) {
return f(request)
}
func TestStreamGeminiImageHashDoesNotNeedWholeResponse(t *testing.T) {
payload := paddedPNG(4 << 20)
response := fmt.Sprintf(`{"candidates":[{"content":{"parts":[{"inlineData":{"mimeType":"image/png","data":"%s"}}]}}]}`,
@@ -44,8 +50,8 @@ func TestPaddedPNGVariantsAreExactSizeAndUnique(t *testing.T) {
}
}
func TestStreamGeminiRequestBodyPreservesInput(t *testing.T) {
input := paddedPNGVariant(2<<20, 37)
func TestStreamGeminiRequestBodyPreservesMultipleInputs(t *testing.T) {
inputs := geminiInputs(3, 2<<20, "multi-image-test", 37)
var body struct {
Contents []struct {
Parts []struct {
@@ -59,24 +65,26 @@ func TestStreamGeminiRequestBodyPreservesInput(t *testing.T) {
ResponseModalities []string `json:"responseModalities"`
} `json:"generationConfig"`
}
stream := streamGeminiRequestBody(input)
stream := streamGeminiRequestBody(inputs)
defer stream.Close()
if err := json.NewDecoder(stream).Decode(&body); err != nil {
t.Fatalf("decode streamed Gemini request: %v", err)
}
if len(body.Contents) != 1 || len(body.Contents[0].Parts) != 2 || body.Contents[0].Parts[1].InlineData == nil {
if len(body.Contents) != 1 || len(body.Contents[0].Parts) != 4 {
t.Fatalf("unexpected Gemini body structure: %+v", body)
}
inlineData := body.Contents[0].Parts[1].InlineData
if inlineData.MIMEType != "image/png" {
t.Fatalf("mime type=%q", inlineData.MIMEType)
}
decoded, err := base64.StdEncoding.DecodeString(inlineData.Data)
if err != nil {
t.Fatalf("decode streamed input: %v", err)
}
if !bytes.Equal(decoded, input) {
t.Fatal("streamed input differs from source")
for index, input := range inputs {
inlineData := body.Contents[0].Parts[index+1].InlineData
if inlineData == nil || inlineData.MIMEType != "image/png" {
t.Fatalf("image %d inline data=%+v", index, inlineData)
}
decoded, err := base64.StdEncoding.DecodeString(inlineData.Data)
if err != nil {
t.Fatalf("decode streamed input %d: %v", index, err)
}
if !bytes.Equal(decoded, input) {
t.Fatalf("streamed input %d differs from source", index)
}
}
if len(body.GenerationConfig.ResponseModalities) != 1 || body.GenerationConfig.ResponseModalities[0] != "IMAGE" {
t.Fatalf("response modalities=%v", body.GenerationConfig.ResponseModalities)
@@ -84,8 +92,8 @@ func TestStreamGeminiRequestBodyPreservesInput(t *testing.T) {
}
func TestGeminiRequestBodyCanBeReplayedAfterHTTP2Failure(t *testing.T) {
input := paddedPNGVariant(256<<10, 91)
request, err := newGeminiRequest(t.Context(), "https://gateway.example/v1beta/models/test:generateContent", input)
inputs := geminiInputs(3, 768<<10, "replay-test", 91)
request, err := newGeminiRequest(t.Context(), "https://gateway.example/v1beta/models/test:generateContent", inputs)
if err != nil {
t.Fatalf("new Gemini request: %v", err)
}
@@ -114,6 +122,37 @@ func TestGeminiRequestBodyCanBeReplayedAfterHTTP2Failure(t *testing.T) {
}
}
func TestGeminiMultiImageRequestUsesThreeFileDataReferences(t *testing.T) {
request, err := newGeminiFileDataRequest(t.Context(), "https://gateway.example/v1beta/models/test:generateContent", []string{
"https://fixtures.example/one.png",
"https://fixtures.example/two.png",
"https://fixtures.example/three.png",
})
if err != nil {
t.Fatalf("new multi-image request: %v", err)
}
var body struct {
Contents []struct {
Parts []struct {
FileData *struct {
FileURI string `json:"fileUri"`
} `json:"fileData"`
} `json:"parts"`
} `json:"contents"`
}
if err := json.NewDecoder(request.Body).Decode(&body); err != nil {
t.Fatalf("decode multi-image request: %v", err)
}
if len(body.Contents) != 1 || len(body.Contents[0].Parts) != 4 {
t.Fatalf("unexpected multi-image body: %+v", body)
}
for index := 1; index < 4; index++ {
if body.Contents[0].Parts[index].FileData == nil || body.Contents[0].Parts[index].FileData.FileURI == "" {
t.Fatalf("missing fileData at part %d", index)
}
}
}
func TestVideoCombinationsProvide128UniqueInputs(t *testing.T) {
images := make([]string, 16)
for index := range images {
@@ -133,6 +172,9 @@ func TestVideoCombinationsProvide128UniqueInputs(t *testing.T) {
if len(seen) != 128 {
t.Fatalf("unique combinations=%d", len(seen))
}
if got := requestedVideoCombinationCount(combinations, 144, 0, 1); got != 131 {
t.Fatalf("requested combinations=%d, want 131", got)
}
}
func TestGeminiLoadIsSplitAcrossTwoGatewayAPIs(t *testing.T) {
@@ -178,7 +220,7 @@ func TestGeminiLoadIsSplitAcrossTwoGatewayAPIs(t *testing.T) {
gateways: []string{first.URL, second.URL}, apiKeys: []string{"key-1", "key-2"}, runID: "run-1",
runToken: "token-1", geminiModel: "gemini-image-test",
}
result := runGemini(t.Context(), http.DefaultClient, opts, "dual-api", 8, 256<<10, 256<<10, false, 0, 1)
result := runGemini(t.Context(), http.DefaultClient, opts, "dual-api", 8, 1, 256<<10, 256<<10, false, 0, 1)
if result.err != nil || result.report.Completed != 8 {
t.Fatalf("result=%+v err=%v", result.report, result.err)
}
@@ -254,6 +296,21 @@ func TestValidateVideoAssetDownloadsFinalMedia(t *testing.T) {
}
}
func TestValidateVideoAssetUsesIndependentTransportForExternalMedia(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "video/mp4")
_, _ = w.Write([]byte{0, 0, 0, 16, 'f', 't', 'y', 'p', 'i', 's', 'o', 'm'})
}))
defer server.Close()
poisoned := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
return nil, errors.New("gateway-only transport must not be used for external media")
})}
opts := options{gatewayTLSName: "gateway.easyai.local"}
if err := validateVideoAsset(t.Context(), poisoned, opts, server.URL+"/result.mp4", 0); err != nil {
t.Fatalf("validate external video with independent transport: %v", err)
}
}
func TestValidateVideoAssetUsesGatewayForMaterializedPath(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
if request.Host != "gateway.easyai.local" || request.Header.Get("Authorization") != "Bearer key-1" {
+126
View File
@@ -2575,6 +2575,30 @@
}
}
},
"/api/admin/runtime/workers": {
"get": {
"security": [
{
"BearerAuth": []
}
],
"produces": [
"application/json"
],
"tags": [
"acceptance"
],
"summary": "获取集群 Worker 实时负载与领取容量",
"responses": {
"200": {
"description": "OK",
"schema": {
"$ref": "#/definitions/store.WorkerClusterRuntime"
}
}
}
}
},
"/api/admin/system/acceptance/capacity-profiles": {
"get": {
"security": [
@@ -16347,6 +16371,108 @@
"$ref": "#/definitions/store.GatewayWalletAccount"
}
}
},
"store.WorkerClusterRuntime": {
"type": "object",
"properties": {
"capturedAt": {
"type": "string"
},
"queue": {
"$ref": "#/definitions/store.WorkerQueueRuntime"
},
"workers": {
"type": "array",
"items": {
"$ref": "#/definitions/store.WorkerInstanceRuntime"
}
}
}
},
"store.WorkerInstanceRuntime": {
"type": "object",
"properties": {
"activeLeases": {
"type": "integer"
},
"allocatedCapacity": {
"type": "integer"
},
"capacityLimit": {
"type": "integer"
},
"drainingAt": {
"type": "string"
},
"finalizingTasks": {
"type": "integer"
},
"hardCapacityLimit": {
"type": "integer"
},
"heartbeatAt": {
"type": "string"
},
"heavyCapacity": {
"type": "integer"
},
"instanceId": {
"type": "string"
},
"loadSampledAt": {
"type": "string"
},
"podName": {
"type": "string"
},
"podUid": {
"type": "string"
},
"preparingTasks": {
"type": "integer"
},
"pressureReason": {
"type": "string"
},
"pressureState": {
"type": "string"
},
"reportedActiveTasks": {
"type": "integer"
},
"revision": {
"type": "string"
},
"runningTasks": {
"type": "integer"
},
"safeCapacity": {
"type": "integer"
},
"site": {
"type": "string"
},
"status": {
"type": "string"
},
"waitingUpstreamTasks": {
"type": "integer"
}
}
},
"store.WorkerQueueRuntime": {
"type": "object",
"properties": {
"oldestWaitSeconds": {
"type": "number"
},
"queued": {
"type": "integer"
},
"running": {
"type": "integer"
}
}
}
},
"securityDefinitions": {
+81
View File
@@ -4196,6 +4196,73 @@ definitions:
primaryAccount:
$ref: '#/definitions/store.GatewayWalletAccount'
type: object
store.WorkerClusterRuntime:
properties:
capturedAt:
type: string
queue:
$ref: '#/definitions/store.WorkerQueueRuntime'
workers:
items:
$ref: '#/definitions/store.WorkerInstanceRuntime'
type: array
type: object
store.WorkerInstanceRuntime:
properties:
activeLeases:
type: integer
allocatedCapacity:
type: integer
capacityLimit:
type: integer
drainingAt:
type: string
finalizingTasks:
type: integer
hardCapacityLimit:
type: integer
heartbeatAt:
type: string
heavyCapacity:
type: integer
instanceId:
type: string
loadSampledAt:
type: string
podName:
type: string
podUid:
type: string
preparingTasks:
type: integer
pressureReason:
type: string
pressureState:
type: string
reportedActiveTasks:
type: integer
revision:
type: string
runningTasks:
type: integer
safeCapacity:
type: integer
site:
type: string
status:
type: string
waitingUpstreamTasks:
type: integer
type: object
store.WorkerQueueRuntime:
properties:
oldestWaitSeconds:
type: number
queued:
type: integer
running:
type: integer
type: object
info:
contact: {}
description: |-
@@ -5846,6 +5913,20 @@ paths:
summary: 更新 Runner 策略
tags:
- runtime
/api/admin/runtime/workers:
get:
produces:
- application/json
responses:
"200":
description: OK
schema:
$ref: '#/definitions/store.WorkerClusterRuntime'
security:
- BearerAuth: []
summary: 获取集群 Worker 实时负载与领取容量
tags:
- acceptance
/api/admin/system/acceptance/capacity-profiles:
get:
produces:
+30 -8
View File
@@ -59,6 +59,7 @@ type Server struct {
type Report struct {
GeminiRequests int64 `json:"geminiRequests"`
GeminiInvalid int64 `json:"geminiInvalid"`
GeminiInputImages int64 `json:"geminiInputImages"`
GeminiInputBytes int64 `json:"geminiInputBytes"`
GeminiOutputBytes int64 `json:"geminiOutputBytes"`
VideoSubmissions int64 `json:"videoSubmissions"`
@@ -163,7 +164,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
writeProtocolError(w, http.StatusBadRequest, err.Error())
return
}
inputBytes, err := validateGeminiImageRequest(body)
inputBytes, inputImages, err := validateGeminiImageRequest(body)
if err != nil {
s.recordGeminiInvalid()
writeProtocolError(w, http.StatusBadRequest, err.Error())
@@ -180,6 +181,7 @@ func (s *Server) geminiGenerateContent(w http.ResponseWriter, r *http.Request) {
s.geminiIdempotency[idempotencyKey] = struct{}{}
}
s.report.GeminiRequests++
s.report.GeminiInputImages += int64(inputImages)
s.report.GeminiInputBytes += int64(inputBytes)
s.report.GeminiOutputBytes += int64(outputBytes)
s.mu.Unlock()
@@ -533,7 +535,7 @@ func decodeImageDataURL(raw string) ([]byte, error) {
return payload, nil
}
func validateGeminiImageRequest(body map[string]any) (int, error) {
func validateGeminiImageRequest(body map[string]any) (int, int, error) {
generationConfig, _ := body["generationConfig"].(map[string]any)
modalities, _ := generationConfig["responseModalities"].([]any)
hasImageModality := false
@@ -543,9 +545,10 @@ func validateGeminiImageRequest(body map[string]any) (int, error) {
}
}
if !hasImageModality {
return 0, errors.New("generationConfig.responseModalities must contain IMAGE")
return 0, 0, errors.New("generationConfig.responseModalities must contain IMAGE")
}
total := 0
images := 0
contents, _ := body["contents"].([]any)
for _, rawContent := range contents {
content, _ := rawContent.(map[string]any)
@@ -558,22 +561,41 @@ func validateGeminiImageRequest(body map[string]any) (int, error) {
}
encoded := strings.TrimSpace(stringValue(inline["data"]))
if encoded == "" {
fileData, _ := part["fileData"].(map[string]any)
if fileData == nil {
fileData, _ = part["file_data"].(map[string]any)
}
fileURI := strings.TrimSpace(stringValue(fileData["fileUri"]))
if fileURI == "" {
fileURI = strings.TrimSpace(stringValue(fileData["file_uri"]))
}
if fileURI == "" {
fileURI = strings.TrimSpace(stringValue(fileData["uri"]))
}
if fileURI == "" {
continue
}
if !strings.HasPrefix(fileURI, "http://") && !strings.HasPrefix(fileURI, "https://") {
return 0, 0, errors.New("Gemini fileData URI must use HTTP or HTTPS")
}
images++
continue
}
decoded, err := base64.StdEncoding.DecodeString(encoded)
if err != nil {
return 0, fmt.Errorf("Gemini inlineData is not valid Base64: %w", err)
return 0, 0, fmt.Errorf("Gemini inlineData is not valid Base64: %w", err)
}
if len(decoded) == 0 {
return 0, errors.New("Gemini inlineData is empty")
return 0, 0, errors.New("Gemini inlineData is empty")
}
total += len(decoded)
images++
}
}
if total == 0 {
return 0, errors.New("Gemini request has no inlineData image")
if images == 0 {
return 0, 0, errors.New("Gemini request has no inlineData or fileData image")
}
return total, nil
return total, images, nil
}
func (s *Server) recordGeminiInvalid() {
@@ -155,6 +155,44 @@ func TestForcedConversionRejectsImagesOutsideOfficialSeedanceRange(t *testing.T)
}
}
func TestGeminiProtocolAcceptsMultipleFileDataImages(t *testing.T) {
server := New(Config{Wait: func(context.Context, time.Duration) error { return nil }})
httpServer := httptest.NewServer(server.Handler())
defer httpServer.Close()
parts := []any{map[string]any{"text": "combine references"}}
for index := 0; index < 3; index++ {
parts = append(parts, map[string]any{"fileData": map[string]any{
"mimeType": "image/png",
"fileUri": fmt.Sprintf("https://fixtures.example/%d.png", index),
}})
}
body, _ := json.Marshal(map[string]any{
"contents": []any{map[string]any{"parts": parts}},
"generationConfig": map[string]any{"responseModalities": []any{"IMAGE"}},
})
response, err := postIdempotent(httpServer.URL+"/v1beta/models/gemini-test:generateContent", body, "gemini-multi")
if err != nil {
t.Fatalf("Gemini multi-image request: %v", err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
payload, _ := io.ReadAll(response.Body)
t.Fatalf("Gemini multi-image status=%d body=%s", response.StatusCode, payload)
}
reportResponse, err := http.Get(httpServer.URL + "/report")
if err != nil {
t.Fatalf("get report: %v", err)
}
defer reportResponse.Body.Close()
var report Report
if err := json.NewDecoder(reportResponse.Body).Decode(&report); err != nil {
t.Fatalf("decode report: %v", err)
}
if report.GeminiRequests != 1 || report.GeminiInputImages != 3 || report.GeminiInvalid != 0 {
t.Fatalf("unexpected multi-image report: %+v", report)
}
}
func TestImageFixturesFullyDecode(t *testing.T) {
for name, fixture := range buildFixtures() {
if !strings.HasPrefix(fixture.ContentType, "image/") {
@@ -5,6 +5,7 @@ import "time"
type GeminiProfile struct {
Name string
Requests int
InputImages int
InputBytes int
OutputBytes int
DelayMin time.Duration
@@ -13,15 +14,19 @@ type GeminiProfile struct {
var (
GeminiBaseline = GeminiProfile{
Name: "gemini-baseline", Requests: 1000, InputBytes: 256 << 10, OutputBytes: 256 << 10,
Name: "gemini-baseline", Requests: 1000, InputImages: 1, InputBytes: 256 << 10, OutputBytes: 256 << 10,
DelayMin: 4 * time.Second, DelayMax: 4 * time.Second,
}
GeminiMultiImage = GeminiProfile{
Name: "gemini-multi-image", Requests: 192, InputImages: 3, InputBytes: 768 << 10, OutputBytes: 0,
DelayMin: 8 * time.Second, DelayMax: 15 * time.Second,
}
GeminiLarge = GeminiProfile{
Name: "gemini-large", Requests: 128, InputBytes: 2 << 20, OutputBytes: 4 << 20,
Name: "gemini-large", Requests: 128, InputImages: 1, InputBytes: 2 << 20, OutputBytes: 4 << 20,
DelayMin: 8 * time.Second, DelayMax: 15 * time.Second,
}
GeminiPeak = GeminiProfile{
Name: "gemini-peak", Requests: 32, InputBytes: 8 << 20, OutputBytes: 8 << 20,
Name: "gemini-peak", Requests: 32, InputImages: 1, InputBytes: 8 << 20, OutputBytes: 8 << 20,
DelayMin: 15 * time.Second, DelayMax: 30 * time.Second,
}
)
@@ -30,6 +35,8 @@ func GeminiProfileByName(name string) (GeminiProfile, bool) {
switch name {
case GeminiBaseline.Name:
return GeminiBaseline, true
case GeminiMultiImage.Name:
return GeminiMultiImage, true
case GeminiLarge.Name:
return GeminiLarge, true
case GeminiPeak.Name:
@@ -3,15 +3,18 @@ package acceptanceworkload
import "testing"
func TestGeminiProfilesAndVideoDistribution(t *testing.T) {
for _, profile := range []GeminiProfile{GeminiBaseline, GeminiLarge, GeminiPeak} {
for _, profile := range []GeminiProfile{GeminiBaseline, GeminiMultiImage, GeminiLarge, GeminiPeak} {
resolved, ok := GeminiProfileByName(profile.Name)
if !ok || resolved != profile {
t.Fatalf("profile %q resolved to %+v, ok=%v", profile.Name, resolved, ok)
}
if profile.Requests <= 0 || profile.InputBytes <= 0 || profile.OutputBytes <= 0 ||
if profile.Requests <= 0 || profile.InputImages <= 0 || profile.InputBytes < profile.InputImages || profile.OutputBytes < 0 ||
profile.DelayMin <= 0 || profile.DelayMax < profile.DelayMin {
t.Fatalf("invalid Gemini profile: %+v", profile)
}
if profile.Name != GeminiMultiImage.Name && profile.OutputBytes == 0 {
t.Fatalf("fixed-output Gemini profile has no output size: %+v", profile)
}
}
counts := map[int]int{}
for index := 0; index < 100; index++ {
+7
View File
@@ -79,6 +79,7 @@ type Config struct {
AsyncWorkerHardLimit int
AsyncWorkerInstanceHardLimit int
AsyncWorkerRefreshIntervalSeconds int
AsyncWorkerLoadMode string
AsyncAdmissionMicrobatchSize int
AsyncAdmissionDispatcherEnabled bool
AsyncAdmissionDispatcherConfigured bool
@@ -189,6 +190,7 @@ func Load() Config {
),
AsyncWorkerInstanceHardLimit: envIntValidated("AI_GATEWAY_ASYNC_WORKER_INSTANCE_HARD_LIMIT", 32),
AsyncWorkerRefreshIntervalSeconds: envIntValidated("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS", 5),
AsyncWorkerLoadMode: strings.ToLower(strings.TrimSpace(env("AI_GATEWAY_WORKER_LOAD_MODE", "adaptive"))),
AsyncAdmissionMicrobatchSize: envIntValidated("AI_GATEWAY_ASYNC_ADMISSION_MICROBATCH_SIZE", 8),
AsyncAdmissionDispatcherEnabled: envValue("AI_GATEWAY_ASYNC_ADMISSION_DISPATCHER_ENABLED") == "true",
AsyncAdmissionDispatcherConfigured: envValue(
@@ -299,6 +301,11 @@ func (c Config) Validate() error {
if c.AsyncWorkerRefreshIntervalSeconds < 1 {
return errors.New("AI_GATEWAY_ASYNC_WORKER_REFRESH_INTERVAL_SECONDS must be positive")
}
switch strings.ToLower(strings.TrimSpace(c.AsyncWorkerLoadMode)) {
case "", "adaptive", "legacy":
default:
return errors.New("AI_GATEWAY_WORKER_LOAD_MODE must be adaptive or legacy")
}
if c.AsyncAdmissionMicrobatchSize < 1 || c.AsyncAdmissionMicrobatchSize > 32 {
return errors.New("AI_GATEWAY_ASYNC_ADMISSION_MICROBATCH_SIZE must be between 1 and 32")
}
+5
View File
@@ -95,6 +95,11 @@ func TestValidateAsyncWorkerSettings(t *testing.T) {
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "REFRESH_INTERVAL") {
t.Fatalf("Validate() error = %v, want invalid refresh interval", err)
}
cfg.AsyncWorkerRefreshIntervalSeconds = 5
cfg.AsyncWorkerLoadMode = "invalid"
if err := cfg.Validate(); err == nil || !strings.Contains(err.Error(), "LOAD_MODE") {
t.Fatalf("Validate() error = %v, want invalid worker load mode", err)
}
t.Setenv("AI_GATEWAY_ASYNC_WORKER_HARD_LIMIT", "not-an-integer")
loaded := Load()
if err := loaded.Validate(); err == nil || !strings.Contains(err.Error(), "HARD_LIMIT") {
@@ -264,6 +264,22 @@ func (s *Server) listCapacityProfiles(w http.ResponseWriter, r *http.Request) {
writeJSON(w, http.StatusOK, profiles)
}
// getWorkerClusterRuntime godoc
// @Summary 获取集群 Worker 实时负载与领取容量
// @Tags acceptance
// @Produce json
// @Security BearerAuth
// @Success 200 {object} store.WorkerClusterRuntime
// @Router /api/admin/runtime/workers [get]
func (s *Server) getWorkerClusterRuntime(w http.ResponseWriter, r *http.Request) {
runtime, err := s.acceptanceStore().GetWorkerClusterRuntime(r.Context())
if err != nil {
writeError(w, http.StatusInternalServerError, "get worker cluster runtime failed")
return
}
writeJSON(w, http.StatusOK, runtime)
}
// createAcceptanceRun godoc
// @Summary 创建生产同构验收 Run
// @Tags acceptance
+1
View File
@@ -264,6 +264,7 @@ func NewServerWithStores(
mux.Handle("POST /api/admin/system/acceptance/runs/{runID}/promote", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.promoteAcceptanceRun)))
mux.Handle("POST /api/admin/system/acceptance/runs/{runID}/abort", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.abortAcceptanceRun)))
mux.Handle("GET /api/admin/system/acceptance/capacity-profiles", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.listCapacityProfiles)))
mux.Handle("GET /api/admin/runtime/workers", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getWorkerClusterRuntime)))
mux.Handle("GET /api/admin/system/identity/configuration", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getIdentityConfiguration)))
mux.Handle("POST /api/admin/system/identity/pairings", server.requireAdmin(auth.PermissionManager, http.HandlerFunc(server.startIdentityPairing)))
mux.Handle("GET /api/admin/system/identity/pairings/{pairingID}", server.requireAdmin(auth.PermissionPower, http.HandlerFunc(server.getIdentityPairing)))
+32 -19
View File
@@ -205,23 +205,12 @@ func (s *Service) taskAdmissionScopes(
}
func acceptanceAdmissionScopes(task store.GatewayTask, scopes []store.AdmissionScope) []store.AdmissionScope {
if task.RunMode != "acceptance" && task.RunMode != "acceptance_canary" {
if task.RunMode != "acceptance" {
return scopes
}
out := append([]store.AdmissionScope(nil), scopes...)
for index := range out {
// Protocol-emulated acceptance measures Gateway and Worker capacity, so
// the isolated Run is bounded by the worker_capacity scope instead of a
// production supplier quota. The real acceptance_canary path deliberately
// retains the production platform-model concurrency limit.
if task.RunMode == "acceptance" && out[index].ScopeType == "platform_model" {
out[index].ConcurrentLimit = 0
}
if out[index].ConcurrentLimit <= 0 {
if task.RunMode != "acceptance" || out[index].ScopeType != "platform_model" {
continue
}
}
out[index].ScopeKey = acceptanceScopeKey(task, out[index].ScopeKey)
out[index].QueueLimit = acceptanceQueueLimit
out[index].MaxWaitSeconds = acceptanceQueueMaxWait
}
@@ -235,16 +224,21 @@ func acceptanceInfrastructureReservations(
if task.RunMode != "acceptance" {
return reservations
}
out := make([]store.RateLimitReservation, 0, len(reservations))
for _, reservation := range reservations {
if reservation.ScopeType == "platform_model" && reservation.Metric == "concurrent" {
continue
}
out = append(out, reservation)
out := append([]store.RateLimitReservation(nil), reservations...)
for index := range out {
out[index].ScopeKey = acceptanceScopeKey(task, out[index].ScopeKey)
}
return out
}
func acceptanceScopeKey(task store.GatewayTask, scopeKey string) string {
runID := strings.TrimSpace(task.AcceptanceRunID)
if runID == "" {
runID = "unbound"
}
return "acceptance:" + runID + ":" + scopeKey
}
func (s *Service) loadAsyncTaskAdmission(ctx context.Context, task store.GatewayTask) (*store.TaskAdmission, error) {
if !task.AsyncMode {
return nil, nil
@@ -265,6 +259,7 @@ func pinCandidatesToTaskAdmission(
) ([]store.RuntimeModelCandidate, bool) {
if admission == nil ||
(admission.Status != "waiting" && admission.Status != "admitted") ||
!admission.ReselectRequestedAt.IsZero() ||
len(candidates) < 2 {
return candidates, false
}
@@ -785,6 +780,24 @@ func (s *Service) dispatchWaitingAsyncTasks(ctx context.Context, admissions []st
return false, outcome.Err
}
if !outcome.Result.Admitted {
platformModelID := outcome.Result.Admission.PlatformModelID
if platformModelID == "" {
for _, input := range inputs {
if input.TaskID == outcome.TaskID {
platformModelID = input.PlatformModelID
break
}
}
}
if platformModelID != "" {
marked, markErr := s.store.RequestWaitingTaskAdmissionReselect(ctx, platformModelID)
if markErr != nil {
return false, markErr
}
if marked > 0 {
s.observeTaskAdmission("candidate_reselect_requested")
}
}
return true, nil
}
}
+33 -11
View File
@@ -57,7 +57,7 @@ func TestDistributedAdmissionModelTypeBoundary(t *testing.T) {
}
}
func TestAcceptanceAdmissionScopesUseWorkerCapacityInsteadOfSupplierConcurrency(t *testing.T) {
func TestAcceptanceAdmissionScopesIsolateRunWithoutChangingLimits(t *testing.T) {
input := []store.AdmissionScope{{
ScopeType: "platform_model",
ScopeKey: "model-1",
@@ -70,15 +70,15 @@ func TestAcceptanceAdmissionScopesUseWorkerCapacityInsteadOfSupplierConcurrency(
ConcurrentLimit: 48,
}}
got := acceptanceAdmissionScopes(store.GatewayTask{RunMode: "acceptance"}, input)
got := acceptanceAdmissionScopes(store.GatewayTask{RunMode: "acceptance", AcceptanceRunID: "run-1"}, input)
if got[0].ConcurrentLimit != 0 {
t.Fatalf("protocol-emulated acceptance must defer to worker capacity, got %+v", got[0])
if got[0].ConcurrentLimit != 10 || got[0].ScopeKey != "acceptance:run-1:model-1" {
t.Fatalf("protocol-emulated acceptance must preserve the limit in an isolated scope, got %+v", got[0])
}
if got[0].QueueLimit != acceptanceQueueLimit || got[0].MaxWaitSeconds != acceptanceQueueMaxWait {
t.Fatalf("acceptance must enable a bounded queue, got %+v", got[0])
}
if got[1].ConcurrentLimit != 48 {
if got[1].ConcurrentLimit != 48 || got[1].ScopeKey != "acceptance:run-1:global" {
t.Fatalf("acceptance worker capacity changed, got %+v", got[1])
}
if input[0].QueueLimit != 0 || input[0].MaxWaitSeconds != 0 {
@@ -100,15 +100,15 @@ func TestAcceptanceCanaryPreservesProductionConcurrency(t *testing.T) {
}
}
func TestAcceptanceInfrastructureReservationsOnlyRemoveSupplierConcurrency(t *testing.T) {
func TestAcceptanceInfrastructureReservationsIsolateEveryMetric(t *testing.T) {
input := []store.RateLimitReservation{
{ScopeType: "platform_model", Metric: "concurrent", Limit: 10},
{ScopeType: "platform_model", Metric: "rpm", Limit: 600},
{ScopeType: "user_group", Metric: "concurrent", Limit: 20},
{ScopeType: "platform_model", ScopeKey: "model-1", Metric: "concurrent", Limit: 10},
{ScopeType: "platform_model", ScopeKey: "model-1", Metric: "rpm", Limit: 600},
{ScopeType: "user_group", ScopeKey: "group-1", Metric: "concurrent", Limit: 20},
}
got := acceptanceInfrastructureReservations(store.GatewayTask{RunMode: "acceptance"}, input)
if len(got) != 2 || got[0].Metric != "rpm" || got[1].ScopeType != "user_group" {
got := acceptanceInfrastructureReservations(store.GatewayTask{RunMode: "acceptance", AcceptanceRunID: "run-1"}, input)
if len(got) != 3 || got[0].ScopeKey != "acceptance:run-1:model-1" || got[1].ScopeKey != "acceptance:run-1:model-1" || got[2].ScopeKey != "acceptance:run-1:group-1" {
t.Fatalf("unexpected acceptance reservations: %+v", got)
}
canary := acceptanceInfrastructureReservations(store.GatewayTask{RunMode: "acceptance_canary"}, input)
@@ -192,6 +192,28 @@ func TestPinCandidatesToTaskAdmissionPreservesWaitingCandidate(t *testing.T) {
}
}
func TestPinCandidatesToTaskAdmissionAllowsRequestedReselection(t *testing.T) {
input := []store.RuntimeModelCandidate{
{PlatformID: "platform-a", PlatformModelID: "model-a"},
{PlatformID: "platform-b", PlatformModelID: "model-b"},
}
admission := &store.TaskAdmission{
Status: "waiting",
PlatformID: "platform-b",
PlatformModelID: "model-b",
ReselectRequestedAt: time.Now(),
}
got, pinned := pinCandidatesToTaskAdmission(input, admission)
if pinned {
t.Fatalf("reselection request unexpectedly pinned the old candidate: %+v", got)
}
if got[0].PlatformModelID != "model-a" {
t.Fatalf("reselection changed the sorted candidate order: %+v", got)
}
}
func TestPinCandidatesToTaskAdmissionIgnoresMissingCandidate(t *testing.T) {
input := []store.RuntimeModelCandidate{
{PlatformID: "platform-a", PlatformModelID: "model-a"},
+103 -9
View File
@@ -10,8 +10,10 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/workerload"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/riverqueue/river"
"github.com/riverqueue/river/riverdriver/riverpgxv5"
"github.com/riverqueue/river/rivermigrate"
@@ -45,6 +47,10 @@ type asyncTaskWorker struct {
service *Service
}
type workerLoadSampler interface {
Sample(databaseConnections, databaseMax int32) workerload.ResourceSample
}
func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs]) error {
task, err := w.service.store.GetTask(ctx, job.Args.TaskID)
if err != nil {
@@ -53,6 +59,12 @@ func (w *asyncTaskWorker) Work(ctx context.Context, job *river.Job[asyncTaskArgs
if task.Status == "succeeded" || task.Status == "failed" || task.Status == "cancelled" {
return nil
}
loadLease, admitted := w.service.tryStartWorkerTask()
if !admitted {
return river.JobSnooze(workerLoadRetryDelay(task.ID))
}
defer loadLease.Release()
ctx = context.WithValue(ctx, workerLoadLeaseContextKey{}, loadLease)
executionToken := uuid.NewString()
result, runErr := w.service.executeWithToken(ctx, task, authUserFromTask(task), nil, executionToken)
if runErr == nil {
@@ -323,21 +335,36 @@ func (s *Service) makeAsyncExecutionClient(capacity int) (asyncExecutionClient,
func (s *Service) loadAsyncWorkerCapacity(ctx context.Context) (store.AsyncWorkerCapacitySnapshot, error) {
if s.asyncCapacityLoader != nil {
return s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit)
snapshot, err := s.asyncCapacityLoader(ctx, s.cfg.AsyncWorkerHardLimit)
if err == nil && s.workerLoad != nil {
s.workerLoad.SetClaimLimit(snapshot.Capacity)
}
return snapshot, err
}
snapshot, err := s.coordinationStore.AsyncWorkerCapacity(ctx, s.cfg.AsyncWorkerHardLimit)
if err != nil {
return store.AsyncWorkerCapacitySnapshot{}, err
}
loadSnapshot := s.sampleWorkerLoad()
allocation, err := s.coordinationStore.RegisterWorkerInstance(ctx, store.WorkerRegistrationInput{
InstanceID: s.workerInstanceID,
PodUID: strings.TrimSpace(os.Getenv("POD_UID")),
PodName: strings.TrimSpace(os.Getenv("POD_NAME")),
Site: strings.TrimSpace(os.Getenv("EASYAI_SITE")),
Revision: strings.TrimSpace(os.Getenv("AI_GATEWAY_REVISION")),
DesiredCapacity: snapshot.Capacity,
CapacityLimit: s.cfg.AsyncWorkerInstanceHardLimit,
HeartbeatStaleAfter: time.Duration(s.cfg.AsyncWorkerRefreshIntervalSeconds) * 6 * time.Second,
InstanceID: s.workerInstanceID,
PodUID: strings.TrimSpace(os.Getenv("POD_UID")),
PodName: strings.TrimSpace(os.Getenv("POD_NAME")),
Site: strings.TrimSpace(os.Getenv("EASYAI_SITE")),
Revision: strings.TrimSpace(os.Getenv("AI_GATEWAY_REVISION")),
DesiredCapacity: snapshot.Capacity,
CapacityLimit: s.cfg.AsyncWorkerInstanceHardLimit,
LoadMode: loadSnapshot.Mode,
SafeCapacity: loadSnapshot.SafeCapacity,
HeavyCapacity: loadSnapshot.HeavyLimit,
ActiveTasks: loadSnapshot.ActiveTasks,
PreparingTasks: loadSnapshot.PreparingTasks,
WaitingUpstreamTasks: loadSnapshot.WaitingUpstreamTasks,
FinalizingTasks: loadSnapshot.FinalizingTasks,
PressureState: string(loadSnapshot.PressureState),
PressureReason: loadSnapshot.PressureReason,
LoadSampledAt: loadSnapshot.SampledAt,
HeartbeatStaleAfter: time.Duration(s.cfg.AsyncWorkerRefreshIntervalSeconds) * 6 * time.Second,
})
if err != nil {
return store.AsyncWorkerCapacitySnapshot{}, err
@@ -346,9 +373,60 @@ func (s *Service) loadAsyncWorkerCapacity(ctx context.Context) (store.AsyncWorke
snapshot.GlobalCapacity = allocation.GlobalAllocated
snapshot.ActiveInstances = allocation.ActiveInstances
snapshot.InstanceID = allocation.InstanceID
loadSnapshot = s.workerLoad.SetClaimLimit(allocation.Allocated)
snapshot.LoadMode = loadSnapshot.Mode
snapshot.LocalSafeCapacity = loadSnapshot.SafeCapacity
snapshot.LocalHeavyCapacity = loadSnapshot.HeavyLimit
snapshot.LocalActiveTasks = loadSnapshot.ActiveTasks
snapshot.LocalPreparingTasks = loadSnapshot.PreparingTasks
snapshot.LocalWaitingTasks = loadSnapshot.WaitingUpstreamTasks
snapshot.LocalFinalizingTasks = loadSnapshot.FinalizingTasks
snapshot.LocalPressureState = string(loadSnapshot.PressureState)
snapshot.LocalPressureReason = loadSnapshot.PressureReason
return snapshot, nil
}
func (s *Service) sampleWorkerLoad() workerload.Snapshot {
if s.workerLoad == nil {
return workerload.Snapshot{Mode: workerload.ModeLegacy, SafeCapacity: s.cfg.AsyncWorkerInstanceHardLimit, HeavyLimit: s.cfg.AsyncWorkerInstanceHardLimit}
}
if s.workerLoadSampler == nil {
return s.workerLoad.Snapshot()
}
var connections int32
var maximum int32
seen := make(map[*pgxpool.Pool]struct{}, 3)
for _, database := range []*store.Store{s.store, s.coordinationStore, s.riverStore} {
if database == nil || database.Pool() == nil {
continue
}
pool := database.Pool()
if _, ok := seen[pool]; ok {
continue
}
seen[pool] = struct{}{}
statistics := pool.Stat()
connections += statistics.AcquiredConns()
maximum += statistics.MaxConns()
}
return s.workerLoad.Observe(s.workerLoadSampler.Sample(connections, maximum))
}
func (s *Service) tryStartWorkerTask() (*workerload.Lease, bool) {
if s.workerLoad == nil {
return nil, true
}
return s.workerLoad.TryStart()
}
func workerLoadRetryDelay(taskID string) time.Duration {
var value uint32
for _, character := range []byte(taskID) {
value = value*33 + uint32(character)
}
return 250*time.Millisecond + time.Duration(value%501)*time.Millisecond
}
func (s *Service) refreshAsyncWorkerCapacity(ctx context.Context) {
ticker := time.NewTicker(time.Duration(s.cfg.AsyncWorkerRefreshIntervalSeconds) * time.Second)
defer ticker.Stop()
@@ -427,6 +505,12 @@ func (s *Service) resizeAsyncWorkerCapacity(ctx context.Context) {
"modelDesired", snapshot.ModelDesired,
"groupDesired", snapshot.GroupDesired,
"capped", snapshot.Capped,
"loadMode", snapshot.LoadMode,
"localSafeCapacity", snapshot.LocalSafeCapacity,
"localHeavyCapacity", snapshot.LocalHeavyCapacity,
"localActiveTasks", snapshot.LocalActiveTasks,
"pressureState", snapshot.LocalPressureState,
"pressureReason", snapshot.LocalPressureReason,
)
if oldClient != nil {
go s.drainAsyncWorkerClient(oldClient)
@@ -493,6 +577,16 @@ func (s *Service) observeAsyncWorkerCapacity(snapshot store.AsyncWorkerCapacityS
if ok {
distributedObserver.SetDistributedWorkerCapacity(snapshot.ActiveInstances, snapshot.GlobalCapacity, snapshot.Capacity)
}
loadObserver, ok := s.billingMetrics.(interface {
SetWorkerLoad(activeLimit, heavyLimit, active, preparing, waiting, finalizing int, pressure string)
})
if ok {
loadObserver.SetWorkerLoad(
snapshot.LocalSafeCapacity, snapshot.LocalHeavyCapacity, snapshot.LocalActiveTasks,
snapshot.LocalPreparingTasks, snapshot.LocalWaitingTasks, snapshot.LocalFinalizingTasks,
snapshot.LocalPressureState,
)
}
}
func (s *Service) observeAsyncWorkerResize(outcome string) {
+43 -7
View File
@@ -19,6 +19,7 @@ import (
"github.com/easyai/easyai-ai-gateway/apps/api/internal/config"
scriptengine "github.com/easyai/easyai-ai-gateway/apps/api/internal/script"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/store"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/workerload"
"github.com/google/uuid"
"github.com/jackc/pgx/v5"
"github.com/riverqueue/river"
@@ -39,6 +40,8 @@ type Service struct {
riverDrainingClients map[asyncExecutionClient]struct{}
riverWorkerCapacity int
workerInstanceID string
workerLoad *workerload.Controller
workerLoadSampler workerLoadSampler
asyncCapacityLoader func(context.Context, int) (store.AsyncWorkerCapacitySnapshot, error)
admissionWakeMu sync.Mutex
admissionWake chan struct{}
@@ -129,6 +132,9 @@ func NewWithStores(
if cfg.AsyncWorkerRefreshIntervalSeconds == 0 {
cfg.AsyncWorkerRefreshIntervalSeconds = 5
}
if strings.TrimSpace(cfg.AsyncWorkerLoadMode) == "" {
cfg.AsyncWorkerLoadMode = workerload.ModeAdaptive
}
if cfg.MediaMaterializationConcurrency == 0 {
cfg.MediaMaterializationConcurrency = 8
}
@@ -184,8 +190,13 @@ func NewWithStores(
"universal": clients.UniversalClient{HTTPClient: httpClients.none, ScriptExecutor: scriptExecutor},
"simulation": clients.SimulationClient{},
},
httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
httpClients: httpClients,
workerInstanceID: asyncWorkerID(),
workerLoad: workerload.New(workerload.Config{
Mode: cfg.AsyncWorkerLoadMode, HardLimit: cfg.AsyncWorkerInstanceHardLimit,
InitialActive: 4, InitialHeavy: 1, HealthySamples: 3,
}),
workerLoadSampler: workerload.NewSystemSampler(),
admissionWake: make(chan struct{}, 4096),
asyncAdmissionWake: make(chan struct{}, 1),
admissionTaskWaiters: map[string]*admissionTaskWaiter{},
@@ -1266,7 +1277,7 @@ func (s *Service) runCandidate(
}
simulated := isSimulation(task, candidate)
baseAttemptMetrics := mergeMetrics(attemptMetrics(candidate, attemptNo, simulated), parameterPreprocessingMetrics(preprocessing))
reservations := s.rateLimitReservations(ctx, user, candidate, body)
reservations := acceptanceInfrastructureReservations(task, s.rateLimitReservations(ctx, user, candidate, body))
if admittedPlatformModelID == candidate.PlatformModelID && len(admittedLeases) > 0 {
filtered := make([]store.RateLimitReservation, 0, len(reservations))
for _, reservation := range reservations {
@@ -1278,6 +1289,10 @@ func (s *Service) runCandidate(
}
limitResult, err := s.store.ReserveRateLimits(ctx, task.ID, "", reservations)
if err != nil {
var limitErr *store.RateLimitExceededError
if errors.As(err, &limitErr) {
s.observeProviderQuotaWait(limitErr.Metric)
}
retryable := store.RateLimitRetryable(err)
clientErr := &clients.ClientError{Code: "rate_limit", Message: err.Error(), Retryable: retryable}
return clients.Response{}, &localRateLimitError{clientErr: clientErr, cause: err, retryAfter: localRateLimitRetryAfter(err)}
@@ -1438,6 +1453,9 @@ func (s *Service) runCandidate(
); err != nil {
return clients.Response{}, fmt.Errorf("restore upstream submission status: %w", err)
}
if err := enterWorkerWaiting(ctx); err != nil {
return clients.Response{}, err
}
}
setSubmissionStatus := func(status string) error {
if submissionStatus == "response_received" && status != "response_received" {
@@ -1484,7 +1502,10 @@ func (s *Service) runCandidate(
if err := s.persistCompatibilitySubmission(context.WithoutCancel(ctx), task, candidate, remoteTaskID, checkpoint, submissionWire); err != nil {
return err
}
return setSubmissionStatus("response_received")
if err := setSubmissionStatus("response_received"); err != nil {
return err
}
return enterWorkerWaiting(ctx)
},
OnRemoteTaskPolled: func(remoteTaskID string, payload map[string]any) error {
if strings.TrimSpace(remoteTaskID) == "" {
@@ -1496,18 +1517,21 @@ func (s *Service) runCandidate(
}
task.RemoteTaskID = remoteTaskID
task.RemoteTaskPayload = checkpoint
return setSubmissionStatus("response_received")
if err := setSubmissionStatus("response_received"); err != nil {
return err
}
return enterWorkerWaiting(ctx)
},
OnUpstreamSubmissionStarted: func() error {
if err := setSubmissionStatus("submitting"); err != nil {
return err
}
markUpstreamSubmissionStarted(ctx)
return nil
return enterWorkerWaiting(ctx)
},
OnUpstreamResponseReceived: func() error {
submissionStatus = "response_received"
return nil
return enterWorkerFinalizing(ctx)
},
OnUpstreamWireResponse: func(wire *clients.WireResponse) error {
submissionWire = wire
@@ -1521,6 +1545,9 @@ func (s *Service) runCandidate(
UpstreamPreviousResponseID: responseExecution.UpstreamPreviousResponseID,
PreviousResponseTurns: responseExecution.PreviousTurns,
})
if phaseErr := enterWorkerFinalizing(runCtx); err == nil && phaseErr != nil {
err = phaseErr
}
if leaseErr := stopLeaseRenewal(); leaseErr != nil {
err = &clients.ClientError{
Code: "concurrency_lease_lost",
@@ -1717,6 +1744,15 @@ func (s *Service) runCandidate(
return response, nil
}
func (s *Service) observeProviderQuotaWait(metric string) {
observer, ok := s.billingMetrics.(interface {
ObserveProviderQuotaWait(string)
})
if ok {
observer.ObserveProviderQuotaWait(metric)
}
}
func minimalRemoteTaskCheckpoint(provider string, specType string, payload map[string]any) map[string]any {
const maxBytes = 8192
provider = strings.ToLower(strings.TrimSpace(provider))
+28
View File
@@ -0,0 +1,28 @@
package runner
import (
"context"
"github.com/easyai/easyai-ai-gateway/apps/api/internal/workerload"
)
type workerLoadLeaseContextKey struct{}
func workerLoadLease(ctx context.Context) *workerload.Lease {
lease, _ := ctx.Value(workerLoadLeaseContextKey{}).(*workerload.Lease)
return lease
}
func enterWorkerWaiting(ctx context.Context) error {
if lease := workerLoadLease(ctx); lease != nil {
return lease.EnterWaiting()
}
return nil
}
func enterWorkerFinalizing(ctx context.Context) error {
if lease := workerLoadLease(ctx); lease != nil {
return lease.EnterFinalizing(ctx)
}
return nil
}
+138 -47
View File
@@ -5,6 +5,7 @@ import (
"errors"
"fmt"
"net/http"
"strings"
"sync/atomic"
"time"
@@ -22,53 +23,64 @@ type DynamicMetricsSnapshotProvider interface {
}
type Metrics struct {
accepted atomic.Uint64
rejected atomic.Uint64
duplicate atomic.Uint64
sessionsDeleted atomic.Uint64
watermarkRejected atomic.Uint64
verificationAccepted atomic.Uint64
heartbeatAccepted atomic.Uint64
heartbeatFailed atomic.Uint64
introspectionActive atomic.Uint64
introspectionInactive atomic.Uint64
introspectionFailed atomic.Uint64
jwksSSFFailed atomic.Uint64
jwksOIDCFailed atomic.Uint64
processingCount atomic.Uint64
processingNanos atomic.Uint64
processingBuckets [6]atomic.Uint64
billingSettlementCompleted atomic.Uint64
billingSettlementRetry atomic.Uint64
billingManualReview atomic.Uint64
billingEstimateFailed atomic.Uint64
billingIdempotentReplay atomic.Uint64
billingPricingUnavailable atomic.Uint64
asyncWorkerCapacity atomic.Int64
asyncWorkerDesiredCapacity atomic.Int64
asyncWorkerHardLimit atomic.Int64
asyncWorkerCapacityCapped atomic.Int64
asyncWorkerActiveInstances atomic.Int64
asyncWorkerGlobalCapacity atomic.Int64
asyncWorkerAllocated atomic.Int64
asyncWorkerResizeSuccess atomic.Uint64
asyncWorkerRefreshFailed atomic.Uint64
asyncWorkerCreateFailed atomic.Uint64
asyncWorkerStartFailed atomic.Uint64
leaseRenewalSuccess atomic.Uint64
leaseRenewalFailure atomic.Uint64
leaseRenewalLost atomic.Uint64
taskEventDuplicate atomic.Uint64
taskEventUnknownType atomic.Uint64
taskEventBudgetExceeded atomic.Uint64
taskAdmissionAdmitted atomic.Uint64
taskAdmissionQueueFull atomic.Uint64
taskAdmissionTimeout atomic.Uint64
taskAdmissionCancelled atomic.Uint64
taskAdmissionExpired atomic.Uint64
taskAdmissionMigrated atomic.Uint64
taskAdmissionWaitBuckets [11]atomic.Uint64
taskAdmissionWaitMicros atomic.Uint64
accepted atomic.Uint64
rejected atomic.Uint64
duplicate atomic.Uint64
sessionsDeleted atomic.Uint64
watermarkRejected atomic.Uint64
verificationAccepted atomic.Uint64
heartbeatAccepted atomic.Uint64
heartbeatFailed atomic.Uint64
introspectionActive atomic.Uint64
introspectionInactive atomic.Uint64
introspectionFailed atomic.Uint64
jwksSSFFailed atomic.Uint64
jwksOIDCFailed atomic.Uint64
processingCount atomic.Uint64
processingNanos atomic.Uint64
processingBuckets [6]atomic.Uint64
billingSettlementCompleted atomic.Uint64
billingSettlementRetry atomic.Uint64
billingManualReview atomic.Uint64
billingEstimateFailed atomic.Uint64
billingIdempotentReplay atomic.Uint64
billingPricingUnavailable atomic.Uint64
asyncWorkerCapacity atomic.Int64
asyncWorkerDesiredCapacity atomic.Int64
asyncWorkerHardLimit atomic.Int64
asyncWorkerCapacityCapped atomic.Int64
asyncWorkerActiveInstances atomic.Int64
asyncWorkerGlobalCapacity atomic.Int64
asyncWorkerAllocated atomic.Int64
workerSafeCapacity atomic.Int64
workerHeavyCapacity atomic.Int64
workerActiveTasks atomic.Int64
workerPreparingTasks atomic.Int64
workerWaitingTasks atomic.Int64
workerFinalizingTasks atomic.Int64
workerPressureState atomic.Int64
providerQuotaWaitRPM atomic.Uint64
providerQuotaWaitTPM atomic.Uint64
providerQuotaWaitConcurrent atomic.Uint64
providerQuotaWaitOther atomic.Uint64
asyncWorkerResizeSuccess atomic.Uint64
asyncWorkerRefreshFailed atomic.Uint64
asyncWorkerCreateFailed atomic.Uint64
asyncWorkerStartFailed atomic.Uint64
leaseRenewalSuccess atomic.Uint64
leaseRenewalFailure atomic.Uint64
leaseRenewalLost atomic.Uint64
taskEventDuplicate atomic.Uint64
taskEventUnknownType atomic.Uint64
taskEventBudgetExceeded atomic.Uint64
taskAdmissionAdmitted atomic.Uint64
taskAdmissionQueueFull atomic.Uint64
taskAdmissionTimeout atomic.Uint64
taskAdmissionCancelled atomic.Uint64
taskAdmissionExpired atomic.Uint64
taskAdmissionMigrated atomic.Uint64
taskAdmissionWaitBuckets [11]atomic.Uint64
taskAdmissionWaitMicros atomic.Uint64
}
var processingDurationBounds = [...]time.Duration{
@@ -178,6 +190,38 @@ func (m *Metrics) SetDistributedWorkerCapacity(activeInstances, globalCapacity,
m.asyncWorkerAllocated.Store(int64(allocatedCapacity))
}
func (m *Metrics) SetWorkerLoad(activeLimit, heavyLimit, active, preparing, waiting, finalizing int, pressure string) {
m.workerSafeCapacity.Store(int64(activeLimit))
m.workerHeavyCapacity.Store(int64(heavyLimit))
m.workerActiveTasks.Store(int64(active))
m.workerPreparingTasks.Store(int64(preparing))
m.workerWaitingTasks.Store(int64(waiting))
m.workerFinalizingTasks.Store(int64(finalizing))
state := int64(-1)
switch pressure {
case "normal":
state = 0
case "busy":
state = 1
case "critical":
state = 2
}
m.workerPressureState.Store(state)
}
func (m *Metrics) ObserveProviderQuotaWait(metric string) {
switch metric {
case "rpm":
m.providerQuotaWaitRPM.Add(1)
case "tpm", "tpm_total", "tpm_input", "tpm_output":
m.providerQuotaWaitTPM.Add(1)
case "concurrent":
m.providerQuotaWaitConcurrent.Add(1)
default:
m.providerQuotaWaitOther.Add(1)
}
}
func (m *Metrics) ObserveTaskAdmission(event string) {
switch event {
case "admitted":
@@ -297,6 +341,17 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
}); ok {
riverPostgresPool = poolProvider.RiverPostgresPoolMetrics()
}
modelRateLimits := []store.ModelRateLimitStatus{}
if rateLimitProvider, ok := provider.(interface {
ListModelRateLimitStatuses(context.Context) ([]store.ModelRateLimitStatus, error)
}); ok {
var err error
modelRateLimits, err = rateLimitProvider.ListModelRateLimitStatuses(r.Context())
if err != nil {
http.Error(w, "metrics unavailable", http.StatusServiceUnavailable)
return
}
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
outcomeCounters(w, "easyai_gateway_ssf_receipts_total", "Received SETs by bounded outcome.", []outcomeValue{
{"accepted", m.accepted.Load()}, {"rejected", m.rejected.Load()}, {"duplicate", m.duplicate.Load()},
@@ -368,6 +423,20 @@ func (m *Metrics) Handler(provider MetricsSnapshotProvider, issuer, audience str
plainGauge(w, "easyai_gateway_worker_active_instances", "Active distributed worker instances with a fresh heartbeat.", m.asyncWorkerActiveInstances.Load())
plainGauge(w, "easyai_gateway_worker_global_capacity", "Global asynchronous execution capacity before instance allocation.", m.asyncWorkerGlobalCapacity.Load())
plainGauge(w, "easyai_gateway_worker_allocated_capacity", "Asynchronous execution capacity allocated to this worker instance.", m.asyncWorkerAllocated.Load())
plainGauge(w, "easyai_gateway_worker_safe_capacity", "Locally resource-safe asynchronous task capacity.", m.workerSafeCapacity.Load())
plainGauge(w, "easyai_gateway_worker_heavy_capacity", "Locally resource-safe preparing and finalizing capacity.", m.workerHeavyCapacity.Load())
plainGauge(w, "easyai_gateway_worker_active_tasks", "Tasks currently owned by this Worker process.", m.workerActiveTasks.Load())
plainGauge(w, "easyai_gateway_worker_preparing_tasks", "Tasks in the local preparing phase.", m.workerPreparingTasks.Load())
plainGauge(w, "easyai_gateway_worker_waiting_upstream_tasks", "Tasks waiting for an upstream result.", m.workerWaitingTasks.Load())
plainGauge(w, "easyai_gateway_worker_finalizing_tasks", "Tasks in the local finalizing phase.", m.workerFinalizingTasks.Load())
plainGauge(w, "easyai_gateway_worker_pressure_state", "Local Worker pressure state: -1 unknown, 0 normal, 1 busy, 2 critical.", m.workerPressureState.Load())
outcomeCounters(w, "easyai_gateway_provider_quota_waits_total", "Tasks delayed before an upstream call by a cluster-wide provider quota.", []outcomeValue{
{"rpm", m.providerQuotaWaitRPM.Load()},
{"tpm", m.providerQuotaWaitTPM.Load()},
{"concurrent", m.providerQuotaWaitConcurrent.Load()},
{"other", m.providerQuotaWaitOther.Load()},
})
platformModelRateLimitUtilizationGauges(w, modelRateLimits)
plainGauge(w, "easyai_gateway_postgres_pool_max_connections", "Maximum PostgreSQL connections in this process pool.", int64(postgresPool.MaxConnections))
plainGauge(w, "easyai_gateway_postgres_pool_total_connections", "Current PostgreSQL connections in this process pool.", int64(postgresPool.TotalConnections))
plainGauge(w, "easyai_gateway_postgres_pool_acquired_connections", "Currently acquired PostgreSQL connections in this process pool.", int64(postgresPool.AcquiredConnections))
@@ -463,6 +532,28 @@ func plainFloatGauge(w http.ResponseWriter, name, help string, value float64) {
fmt.Fprintf(w, "# HELP %s %s\n# TYPE %s gauge\n%s %.6f\n", name, help, name, name, value)
}
func platformModelRateLimitUtilizationGauges(w http.ResponseWriter, statuses []store.ModelRateLimitStatus) {
const name = "easyai_gateway_platform_model_rate_limit_utilization"
fmt.Fprintf(w, "# HELP %s Cluster-wide platform model quota utilization.\n# TYPE %s gauge\n", name, name)
for _, status := range statuses {
platformModelID := escapePrometheusLabel(status.PlatformModelID)
for _, metric := range []struct {
name string
ratio float64
}{
{name: "rpm", ratio: status.RPM.Ratio},
{name: "tpm", ratio: status.TPM.Ratio},
{name: "concurrent", ratio: status.Concurrent.Ratio},
} {
fmt.Fprintf(w, "%s{platform_model_id=\"%s\",metric=\"%s\"} %.6f\n", name, platformModelID, metric.name, metric.ratio)
}
}
}
func escapePrometheusLabel(value string) string {
return strings.NewReplacer("\\", "\\\\", "\n", "\\n", "\"", "\\\"").Replace(value)
}
func taskAdmissionWaitHistogram(w http.ResponseWriter, metrics *Metrics) {
const name = "easyai_gateway_task_admission_wait_seconds"
fmt.Fprintf(w, "# HELP %s Time spent waiting for persistent task admission.\n# TYPE %s histogram\n", name, name)
@@ -25,6 +25,15 @@ func (m metricsSnapshot) PostgresPoolMetrics() store.PostgresPoolMetricsSnapshot
return m.pool
}
func (m metricsSnapshot) ListModelRateLimitStatuses(context.Context) ([]store.ModelRateLimitStatus, error) {
return []store.ModelRateLimitStatus{{
PlatformModelID: "platform-model-1",
RPM: store.RateLimitMetricStatus{Ratio: .25},
TPM: store.RateLimitMetricStatus{Ratio: .5},
Concurrent: store.RateLimitMetricStatus{Ratio: .75},
}}, nil
}
func TestMetricsExposeBoundedOutcomesAndState(t *testing.T) {
metrics := &Metrics{}
metrics.ObserveReceipt("accepted", 20*time.Millisecond, 2)
@@ -34,6 +43,9 @@ func TestMetricsExposeBoundedOutcomesAndState(t *testing.T) {
metrics.ObserveIntrospection("failed")
metrics.ObserveJWKSRefreshFailure("ssf")
metrics.SetAsyncWorkerCapacity(96, 128, 96, true)
metrics.SetWorkerLoad(12, 3, 8, 2, 5, 1, "busy")
metrics.ObserveProviderQuotaWait("rpm")
metrics.ObserveProviderQuotaWait("tpm_total")
metrics.ObserveAsyncWorkerResize("success")
metrics.ObserveConcurrencyLeaseRenewal("success")
metrics.ObserveConcurrencyLeaseRenewal("lost")
@@ -65,6 +77,16 @@ func TestMetricsExposeBoundedOutcomesAndState(t *testing.T) {
`easyai_gateway_async_worker_capacity 96`,
`easyai_gateway_async_worker_desired_capacity 128`,
`easyai_gateway_async_worker_capacity_capped 1`,
`easyai_gateway_worker_safe_capacity 12`,
`easyai_gateway_worker_heavy_capacity 3`,
`easyai_gateway_worker_active_tasks 8`,
`easyai_gateway_worker_preparing_tasks 2`,
`easyai_gateway_worker_waiting_upstream_tasks 5`,
`easyai_gateway_worker_finalizing_tasks 1`,
`easyai_gateway_worker_pressure_state 1`,
`easyai_gateway_provider_quota_waits_total{outcome="rpm"} 1`,
`easyai_gateway_provider_quota_waits_total{outcome="tpm"} 1`,
`easyai_gateway_platform_model_rate_limit_utilization{platform_model_id="platform-model-1",metric="concurrent"} 0.750000`,
`easyai_gateway_async_worker_resizes_total{outcome="success"} 1`,
`easyai_gateway_concurrency_lease_renewals_total{outcome="success"} 1`,
`easyai_gateway_concurrency_lease_renewals_total{outcome="lost"} 1`,
+1 -1
View File
@@ -689,7 +689,7 @@ ON CONFLICT (task_id, event_type) DO NOTHING`, runID); err != nil {
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases lease
SET released_at = now()
SET released_at = statement_timestamp()
FROM gateway_tasks task
WHERE lease.task_id = task.id
AND lease.released_at IS NULL
+31 -8
View File
@@ -747,14 +747,14 @@ SELECT COUNT(*)
FROM gateway_concurrency_leases
WHERE task_id = $1::uuid
AND released_at IS NULL
AND expires_at > now()`, taskID).Scan(&count)
AND expires_at > statement_timestamp()`, taskID).Scan(&count)
return count, err
}
func resetTaskAdmissionToWaitingTx(ctx context.Context, tx pgx.Tx, input TaskAdmissionInput) (TaskAdmission, error) {
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, input.TaskID); err != nil {
return TaskAdmission{}, err
@@ -862,7 +862,7 @@ func (s *Store) deleteTaskAdmissionOnce(ctx context.Context, taskID string) erro
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
@@ -884,7 +884,7 @@ SELECT id::text,
FROM gateway_concurrency_leases
WHERE task_id = $1::uuid
AND released_at IS NULL
AND expires_at > now()
AND expires_at > statement_timestamp()
ORDER BY scope_type, scope_key, id`, taskID)
if err != nil {
return nil, err
@@ -924,7 +924,7 @@ WHERE admission.mode = 'async'
FROM gateway_concurrency_leases lease
WHERE lease.task_id = admission.task_id
AND lease.released_at IS NULL
AND lease.expires_at > now()
AND lease.expires_at > statement_timestamp()
)
)
)
@@ -980,7 +980,7 @@ WHERE admission.mode = 'async'
FROM gateway_concurrency_leases lease
WHERE lease.task_id = admission.task_id
AND lease.released_at IS NULL
AND lease.expires_at > now()
AND lease.expires_at > statement_timestamp()
)
)
)
@@ -1006,6 +1006,29 @@ LIMIT $1`, limit)
return admissions, rows.Err()
}
// RequestWaitingTaskAdmissionReselect marks every queued task bound to a
// saturated platform model so the dispatcher can route the next batch to a
// different eligible candidate. The durable marker lets multiple dispatchers
// observe the same decision without assigning tasks to a specific Worker.
func (s *Store) RequestWaitingTaskAdmissionReselect(ctx context.Context, platformModelID string) (int64, error) {
result, err := s.pool.Exec(ctx, `
UPDATE gateway_task_admissions admission
SET reselect_requested_at = now(),
updated_at = now()
FROM gateway_tasks task
WHERE admission.task_id = task.id
AND admission.platform_model_id = $1::uuid
AND admission.mode = 'async'
AND admission.status = 'waiting'
AND task.status = 'queued'
AND task.next_run_at <= now()
AND admission.reselect_requested_at IS NULL`, platformModelID)
if err != nil {
return 0, err
}
return result.RowsAffected(), nil
}
// ListWaitingTaskAdmissionIDs returns one FIFO leader for every independent
// platform-model queue, additionally requiring the task to lead its user-group
// queue when one exists. It is used after capacity is released so API
@@ -1268,12 +1291,12 @@ func admissionScopeStatesTx(ctx context.Context, tx pgx.Tx, scopes []AdmissionSc
if scope.ConcurrentLimit > 0 {
if err := tx.QueryRow(ctx, `
SELECT COALESCE(SUM(lease_value), 0)::float8,
COALESCE(MIN(expires_at), now() + interval '1 second')
COALESCE(MIN(expires_at), statement_timestamp() + interval '1 second')
FROM gateway_concurrency_leases
WHERE scope_type = $1
AND scope_key = $2
AND released_at IS NULL
AND expires_at > now()`, scope.ScopeType, scope.ScopeKey).Scan(&state.Active, &state.NextLeaseExpiration); err != nil {
AND expires_at > statement_timestamp()`, scope.ScopeType, scope.ScopeKey).Scan(&state.Active, &state.NextLeaseExpiration); err != nil {
return nil, err
}
state.Saturated = state.Active+scope.Amount > scope.ConcurrentLimit
@@ -451,6 +451,20 @@ WHERE id = $1::uuid`, queuedAtomicTask.ID, queuedSyntheticRiverJobID)
if listedSnapshot == nil || len(listedSnapshot.Scopes) != len(queuedAdmission.Scopes) {
t.Fatalf("listed admission snapshot=%+v, want %d scopes", listedSnapshot, len(queuedAdmission.Scopes))
}
markedForReselect, err := first.RequestWaitingTaskAdmissionReselect(ctx, platformModelID)
if err != nil {
t.Fatalf("request waiting admission reselection: %v", err)
}
if markedForReselect < 1 {
t.Fatalf("marked admissions=%d, want at least the queued task", markedForReselect)
}
reselectAdmission, err := first.GetTaskAdmission(ctx, queuedAtomicTask.ID)
if err != nil {
t.Fatalf("read admission reselection marker: %v", err)
}
if reselectAdmission.ReselectRequestedAt.IsZero() {
t.Fatal("queued admission was not marked for candidate reselection")
}
var riverJobID int64
if err := first.pool.QueryRow(ctx, `
SELECT
@@ -1296,4 +1310,35 @@ WHERE instance_id = $1`, secondID, (workerHeartbeatStaleAfter + time.Second).Str
if err != nil || second.Allocated != 2 || second.GlobalAllocated != 5 || second.ActiveInstances != 2 {
t.Fatalf("bounded two-worker allocation = %+v, err=%v", second, err)
}
first, err = db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{
InstanceID: firstID, DesiredCapacity: 100, CapacityLimit: 10,
LoadMode: "adaptive", SafeCapacity: 0, HeavyCapacity: 1,
ActiveTasks: 2, WaitingUpstreamTasks: 2,
PressureState: "critical", PressureReason: "memory",
})
if err != nil || first.Allocated != 0 {
t.Fatalf("critical worker allocation = %+v, err=%v", first, err)
}
second, err = db.RegisterWorkerInstance(ctx, WorkerRegistrationInput{
InstanceID: secondID, DesiredCapacity: 100, CapacityLimit: 10,
LoadMode: "adaptive", SafeCapacity: 5, HeavyCapacity: 2,
ActiveTasks: 3, PreparingTasks: 1, WaitingUpstreamTasks: 1, FinalizingTasks: 1,
PressureState: "normal",
})
if err != nil || second.Allocated != 5 || second.GlobalAllocated != 5 {
t.Fatalf("adaptive redistribution allocation = %+v, err=%v", second, err)
}
instances, err := db.ListWorkerInstanceRuntime(ctx)
if err != nil {
t.Fatalf("list adaptive worker runtime: %v", err)
}
foundCritical := false
for _, instance := range instances {
if instance.InstanceID == firstID {
foundCritical = instance.SafeCapacity == 0 && instance.ReportedActiveTasks == 2 && instance.WaitingUpstreamTasks == 2 && instance.PressureState == "critical"
}
}
if !foundCritical {
t.Fatalf("adaptive runtime did not expose the critical worker: %+v", instances)
}
}
@@ -6,19 +6,28 @@ import (
)
type AsyncWorkerCapacitySnapshot struct {
Capacity int
GlobalCapacity int
Desired int
HardLimit int
Capped bool
EnabledModels int
UnlimitedModels int
EnabledGroups int
UnlimitedGroups int
ModelDesired int
GroupDesired int
ActiveInstances int
InstanceID string
Capacity int
GlobalCapacity int
Desired int
HardLimit int
Capped bool
EnabledModels int
UnlimitedModels int
EnabledGroups int
UnlimitedGroups int
ModelDesired int
GroupDesired int
ActiveInstances int
InstanceID string
LoadMode string
LocalSafeCapacity int
LocalHeavyCapacity int
LocalActiveTasks int
LocalPreparingTasks int
LocalWaitingTasks int
LocalFinalizingTasks int
LocalPressureState string
LocalPressureReason string
}
func (s *Store) AsyncWorkerCapacity(ctx context.Context, hardLimit int) (AsyncWorkerCapacitySnapshot, error) {
+1 -1
View File
@@ -100,7 +100,7 @@ LEFT JOIN (
FROM gateway_concurrency_leases
WHERE scope_type = 'platform_model'
AND released_at IS NULL
AND expires_at > now()
AND expires_at > statement_timestamp()
GROUP BY scope_key
) con ON con.scope_key = m.id::text
LEFT JOIN (
+1 -1
View File
@@ -169,7 +169,7 @@ LEFT JOIN (
FROM gateway_concurrency_leases
WHERE scope_type = 'platform_model'
AND released_at IS NULL
AND expires_at > now()
AND expires_at > statement_timestamp()
GROUP BY scope_key
) con ON con.scope_key = m.id::text
LEFT JOIN (
+14 -8
View File
@@ -184,12 +184,12 @@ func reserveConcurrencyLease(ctx context.Context, tx pgx.Tx, taskID string, atte
var nextAvailableAt time.Time
if err := tx.QueryRow(ctx, `
SELECT COALESCE(SUM(lease_value), 0)::float8,
COALESCE(MIN(expires_at), now() + ($3::int * interval '1 second'))
COALESCE(MIN(expires_at), statement_timestamp() + ($3::int * interval '1 second'))
FROM gateway_concurrency_leases
WHERE scope_type = $1
AND scope_key = $2
AND released_at IS NULL
AND expires_at > now()`,
AND expires_at > statement_timestamp()`,
reservation.ScopeType,
reservation.ScopeKey,
reservation.LeaseTTLSeconds,
@@ -218,14 +218,20 @@ WHERE scope_type = $1
}
var leaseID string
if err := tx.QueryRow(ctx, `
INSERT INTO gateway_concurrency_leases (task_id, attempt_id, scope_type, scope_key, lease_value, expires_at)
VALUES ($1::uuid, NULLIF($2, '')::uuid, $3, $4, $5, now() + ($6::int * interval '1 second'))
INSERT INTO gateway_concurrency_leases (
task_id, attempt_id, scope_type, scope_key, lease_value, limit_value, acquired_at, expires_at
)
VALUES (
$1::uuid, NULLIF($2, '')::uuid, $3, $4, $5, $6,
statement_timestamp(), statement_timestamp() + ($7::int * interval '1 second')
)
RETURNING id::text`,
taskID,
attemptID,
reservation.ScopeType,
reservation.ScopeKey,
reservation.Amount,
reservation.Limit,
reservation.LeaseTTLSeconds,
).Scan(&leaseID); err != nil {
return ConcurrencyLease{}, err
@@ -381,7 +387,7 @@ func (s *Store) ReleaseConcurrencyLeases(ctx context.Context, leases []Concurren
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE id = ANY($1::uuid[]) AND released_at IS NULL`, leaseIDs); err != nil {
return err
}
@@ -411,7 +417,7 @@ func (s *Store) RenewConcurrencyLeases(ctx context.Context, leases []Concurrency
}
tag, err := s.pool.Exec(ctx, `
UPDATE gateway_concurrency_leases lease
SET expires_at = now() + (renewal.ttl_seconds * interval '1 second')
SET expires_at = statement_timestamp() + (renewal.ttl_seconds * interval '1 second')
FROM unnest($1::uuid[], $2::int[]) AS renewal(id, ttl_seconds)
WHERE lease.id = renewal.id
AND lease.released_at IS NULL
@@ -705,7 +711,7 @@ WHERE attempt.task_id = task.id
result.FailedAttempts = tag.RowsAffected()
tag, err = tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = ANY($1::uuid[])
AND released_at IS NULL`, taskIDs)
if err != nil {
@@ -772,7 +778,7 @@ FOR UPDATE OF task SKIP LOCKED`, runtimeRecoveryBatchSize)
var result RuntimeRecoveryResult
tag, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = ANY($1::uuid[])
AND released_at IS NULL`, taskIDs)
if err != nil {
@@ -98,23 +98,168 @@ WHERE scope_type = 'platform_model'
}
var active int64
var storedLimit float64
if err := first.Pool().QueryRow(ctx, `
SELECT COUNT(*)
SELECT COUNT(*), COALESCE(MAX(limit_value), 0)::float8
FROM gateway_concurrency_leases
WHERE scope_type = 'platform_model'
AND scope_key = $1
AND released_at IS NULL
AND expires_at > now()`, scopeKey).Scan(&active); err != nil {
AND expires_at > now()`, scopeKey).Scan(&active, &storedLimit); err != nil {
t.Fatalf("count active leases: %v", err)
}
if successes.Load() != 64 || active != 64 {
t.Fatalf("successful reservations=%d active leases=%d, want exactly 64", successes.Load(), active)
if successes.Load() != 64 || active != 64 || storedLimit != 64 {
t.Fatalf("successful reservations=%d active leases=%d stored limit=%.0f, want exactly 64", successes.Load(), active, storedLimit)
}
if peak.Load() > 64 {
t.Fatalf("active lease peak=%d, want <=64", peak.Load())
}
}
func TestConcurrencyLeaseTimestampStartsAtReservationStatement(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run concurrency lease PostgreSQL integration tests")
}
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
database, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect store: %v", err)
}
defer database.Close()
scopeKey := "statement-clock-" + time.Now().UTC().Format("20060102150405.000000000")
taskIDs := createLeaseTestTasks(t, ctx, database, 1, scopeKey)
defer deleteLeaseTestTasks(t, database, taskIDs)
tx, err := database.pool.Begin(ctx)
if err != nil {
t.Fatalf("begin reservation transaction: %v", err)
}
defer rollbackTransaction(tx)
var transactionStartedAt time.Time
if err := tx.QueryRow(ctx, `SELECT now()`).Scan(&transactionStartedAt); err != nil {
t.Fatalf("read transaction start: %v", err)
}
time.Sleep(1100 * time.Millisecond)
lease, err := reserveConcurrencyLease(ctx, tx, taskIDs[0], "", RateLimitReservation{
ScopeType: "platform_model",
ScopeKey: scopeKey,
Metric: "concurrent",
Limit: 1,
Amount: 1,
LeaseTTLSeconds: 30,
})
if err != nil {
t.Fatalf("reserve concurrency lease: %v", err)
}
if err := tx.Commit(ctx); err != nil {
t.Fatalf("commit reservation transaction: %v", err)
}
var acquiredAt, expiresAt time.Time
if err := database.pool.QueryRow(ctx, `
SELECT acquired_at, expires_at
FROM gateway_concurrency_leases
WHERE id = $1::uuid`, lease.ID).Scan(&acquiredAt, &expiresAt); err != nil {
t.Fatalf("read lease timestamps: %v", err)
}
if elapsed := acquiredAt.Sub(transactionStartedAt); elapsed < time.Second {
t.Fatalf("lease acquired_at advanced by %s, want at least 1s after transaction start", elapsed)
}
if ttl := expiresAt.Sub(acquiredAt); ttl != 30*time.Second {
t.Fatalf("lease ttl=%s, want 30s", ttl)
}
}
func TestCounterWindowReservationIsAtomicAcrossPools(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run rate limit PostgreSQL integration tests")
}
ctx, cancel := context.WithTimeout(context.Background(), 45*time.Second)
defer cancel()
first, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect first store: %v", err)
}
defer first.Close()
second, err := Connect(ctx, databaseURL)
if err != nil {
t.Fatalf("connect second store: %v", err)
}
defer second.Close()
tests := []struct {
metric string
limit float64
amount float64
wantSuccesses int64
}{
{metric: "rpm", limit: 37, amount: 1, wantSuccesses: 37},
{metric: "tpm_total", limit: 1_000, amount: 25, wantSuccesses: 40},
}
for _, test := range tests {
t.Run(test.metric, func(t *testing.T) {
scopeKey := "atomic-" + test.metric + "-" + time.Now().UTC().Format("20060102150405.000000000")
taskIDs := createLeaseTestTasks(t, ctx, first, 128, scopeKey)
defer deleteLeaseTestTasks(t, first, taskIDs)
var successes atomic.Int64
var wg sync.WaitGroup
errs := make(chan error, len(taskIDs))
for index, taskID := range taskIDs {
wg.Add(1)
go func(index int, taskID string) {
defer wg.Done()
target := first
if index%2 == 1 {
target = second
}
_, err := target.ReserveRateLimits(ctx, taskID, "", []RateLimitReservation{{
ScopeType: "platform_model",
ScopeKey: scopeKey,
Metric: test.metric,
Limit: test.limit,
Amount: test.amount,
WindowSeconds: 3600,
}})
if err == nil {
successes.Add(1)
return
}
if !errors.Is(err, ErrRateLimited) {
errs <- err
}
}(index, taskID)
}
wg.Wait()
close(errs)
for err := range errs {
t.Fatalf("unexpected reservation error: %v", err)
}
var current float64
if err := first.Pool().QueryRow(ctx, `
SELECT COALESCE(MAX(used_value + reserved_value), 0)::float8
FROM gateway_rate_limit_counters
WHERE scope_type = 'platform_model'
AND scope_key = $1
AND metric = $2`, scopeKey, test.metric).Scan(&current); err != nil {
t.Fatalf("read %s counter: %v", test.metric, err)
}
if successes.Load() != test.wantSuccesses {
t.Fatalf("successful %s reservations=%d, want exactly %d", test.metric, successes.Load(), test.wantSuccesses)
}
wantCurrent := float64(test.wantSuccesses) * test.amount
if current != wantCurrent || current > test.limit {
t.Fatalf("%s current=%.0f, want %.0f and <= %.0f", test.metric, current, wantCurrent, test.limit)
}
})
}
}
func TestConcurrencyLeaseRenewalExtendsAndReleases(t *testing.T) {
databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL"))
if databaseURL == "" {
+5 -5
View File
@@ -757,7 +757,7 @@ WHERE task_id = $1::uuid
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
@@ -922,7 +922,7 @@ RETURNING `+gatewayTaskColumns, taskID, message))
changed = true
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
@@ -1011,7 +1011,7 @@ WHERE task_id = $1::uuid
}
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, taskID); err != nil {
return err
@@ -1720,7 +1720,7 @@ ON CONFLICT (task_id, event_type) DO NOTHING`,
if input.FinalizeAdmission {
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, input.TaskID); err != nil {
return err
@@ -2096,7 +2096,7 @@ ON CONFLICT (task_id, event_type) DO NOTHING`, input.TaskID, string(payloadJSON)
if input.FinalizeAdmission {
if _, err := tx.Exec(ctx, `
UPDATE gateway_concurrency_leases
SET released_at = now()
SET released_at = statement_timestamp()
WHERE task_id = $1::uuid
AND released_at IS NULL`, input.TaskID); err != nil {
return err
+131 -27
View File
@@ -12,14 +12,24 @@ import (
const workerHeartbeatStaleAfter = 30 * time.Second
type WorkerRegistrationInput struct {
InstanceID string
PodUID string
PodName string
Site string
Revision string
DesiredCapacity int
CapacityLimit int
HeartbeatStaleAfter time.Duration
InstanceID string
PodUID string
PodName string
Site string
Revision string
DesiredCapacity int
CapacityLimit int
LoadMode string
SafeCapacity int
HeavyCapacity int
ActiveTasks int
PreparingTasks int
WaitingUpstreamTasks int
FinalizingTasks int
PressureState string
PressureReason string
LoadSampledAt time.Time
HeartbeatStaleAfter time.Duration
}
type WorkerAllocation struct {
@@ -59,9 +69,28 @@ func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistra
if input.CapacityLimit < 0 {
return WorkerAllocation{}, errors.New("worker capacity limit cannot be negative")
}
if input.SafeCapacity < 0 || input.HeavyCapacity < 0 || input.ActiveTasks < 0 || input.PreparingTasks < 0 || input.WaitingUpstreamTasks < 0 || input.FinalizingTasks < 0 {
return WorkerAllocation{}, errors.New("worker load values cannot be negative")
}
if input.ActiveTasks != input.PreparingTasks+input.WaitingUpstreamTasks+input.FinalizingTasks {
return WorkerAllocation{}, errors.New("worker active task count must equal phase task counts")
}
if input.CapacityLimit == 0 {
input.CapacityLimit = input.DesiredCapacity
}
hardCapacityLimit := input.CapacityLimit
if strings.EqualFold(strings.TrimSpace(input.LoadMode), "adaptive") {
input.CapacityLimit = min(input.CapacityLimit, input.SafeCapacity)
}
pressureState := strings.ToLower(strings.TrimSpace(input.PressureState))
switch pressureState {
case "normal", "busy", "critical":
default:
pressureState = "unknown"
}
if input.LoadSampledAt.IsZero() {
input.LoadSampledAt = time.Now()
}
staleAfter := input.HeartbeatStaleAfter
if staleAfter < workerHeartbeatStaleAfter {
staleAfter = workerHeartbeatStaleAfter
@@ -83,9 +112,16 @@ func (s *Store) RegisterWorkerInstance(ctx context.Context, input WorkerRegistra
if _, err := tx.Exec(ctx, `
INSERT INTO gateway_worker_instances (
instance_id, pod_uid, pod_name, site, revision, status,
desired_capacity, capacity_limit, allocated_capacity, started_at, heartbeat_at, updated_at
desired_capacity, capacity_limit, hard_capacity_limit, safe_capacity, heavy_capacity,
active_tasks, preparing_tasks, waiting_upstream_tasks, finalizing_tasks,
pressure_state, pressure_reason, load_sampled_at,
allocated_capacity, started_at, heartbeat_at, updated_at
)
VALUES (
$1, $2, $3, $4, $5, 'active', $6, $7, $8, $9, $10,
$11, $12, $13, $14, $15, $16, $17,
0, now(), now(), now()
)
VALUES ($1, $2, $3, $4, $5, 'active', $6, $7, 0, now(), now(), now())
ON CONFLICT (instance_id) DO UPDATE
SET pod_uid = EXCLUDED.pod_uid,
pod_name = EXCLUDED.pod_name,
@@ -97,6 +133,16 @@ SET pod_uid = EXCLUDED.pod_uid,
END,
desired_capacity = EXCLUDED.desired_capacity,
capacity_limit = EXCLUDED.capacity_limit,
hard_capacity_limit = EXCLUDED.hard_capacity_limit,
safe_capacity = EXCLUDED.safe_capacity,
heavy_capacity = EXCLUDED.heavy_capacity,
active_tasks = EXCLUDED.active_tasks,
preparing_tasks = EXCLUDED.preparing_tasks,
waiting_upstream_tasks = EXCLUDED.waiting_upstream_tasks,
finalizing_tasks = EXCLUDED.finalizing_tasks,
pressure_state = EXCLUDED.pressure_state,
pressure_reason = EXCLUDED.pressure_reason,
load_sampled_at = EXCLUDED.load_sampled_at,
heartbeat_at = now(),
updated_at = now()`,
input.InstanceID,
@@ -106,6 +152,16 @@ SET pod_uid = EXCLUDED.pod_uid,
strings.TrimSpace(input.Revision),
input.DesiredCapacity,
input.CapacityLimit,
hardCapacityLimit,
input.SafeCapacity,
input.HeavyCapacity,
input.ActiveTasks,
input.PreparingTasks,
input.WaitingUpstreamTasks,
input.FinalizingTasks,
pressureState,
strings.TrimSpace(input.PressureReason),
input.LoadSampledAt,
); err != nil {
return WorkerAllocation{}, err
}
@@ -198,7 +254,7 @@ func allocateWorkerCapacities(workers []activeWorkerCapacity, desired int) (map[
for _, worker := range workers {
limit := worker.CapacityLimit
if limit <= 0 {
limit = desired
continue
}
if allocations[worker.InstanceID] >= limit {
continue
@@ -254,24 +310,52 @@ WHERE instance_id = $1
}
type WorkerInstanceRuntime struct {
InstanceID string
PodUID string
PodName string
Site string
Revision string
Status string
Allocated int
CapacityLimit int
RunningTasks int
ActiveLeases int
HeartbeatAt time.Time
DrainingAt *time.Time
InstanceID string `json:"instanceId"`
PodUID string `json:"podUid,omitempty"`
PodName string `json:"podName,omitempty"`
Site string `json:"site,omitempty"`
Revision string `json:"revision,omitempty"`
Status string `json:"status"`
Allocated int `json:"allocatedCapacity"`
CapacityLimit int `json:"capacityLimit"`
HardCapacityLimit int `json:"hardCapacityLimit"`
SafeCapacity int `json:"safeCapacity"`
HeavyCapacity int `json:"heavyCapacity"`
ReportedActiveTasks int `json:"reportedActiveTasks"`
PreparingTasks int `json:"preparingTasks"`
WaitingUpstreamTasks int `json:"waitingUpstreamTasks"`
FinalizingTasks int `json:"finalizingTasks"`
PressureState string `json:"pressureState"`
PressureReason string `json:"pressureReason,omitempty"`
LoadSampledAt *time.Time `json:"loadSampledAt,omitempty"`
RunningTasks int `json:"runningTasks"`
ActiveLeases int `json:"activeLeases"`
HeartbeatAt time.Time `json:"heartbeatAt"`
DrainingAt *time.Time `json:"drainingAt,omitempty"`
}
type WorkerQueueRuntime struct {
Queued int
Running int
OldestWaitSeconds float64
Queued int `json:"queued"`
Running int `json:"running"`
OldestWaitSeconds float64 `json:"oldestWaitSeconds"`
}
type WorkerClusterRuntime struct {
Workers []WorkerInstanceRuntime `json:"workers"`
Queue WorkerQueueRuntime `json:"queue"`
CapturedAt time.Time `json:"capturedAt"`
}
func (s *Store) GetWorkerClusterRuntime(ctx context.Context) (WorkerClusterRuntime, error) {
workers, err := s.ListWorkerInstanceRuntime(ctx)
if err != nil {
return WorkerClusterRuntime{}, err
}
queue, err := s.WorkerQueueRuntime(ctx)
if err != nil {
return WorkerClusterRuntime{}, err
}
return WorkerClusterRuntime{Workers: workers, Queue: queue, CapturedAt: time.Now()}, nil
}
type CapacityDatabaseHealth struct {
@@ -306,6 +390,16 @@ SELECT worker.instance_id,
worker.status,
worker.allocated_capacity,
worker.capacity_limit,
worker.hard_capacity_limit,
worker.safe_capacity,
worker.heavy_capacity,
worker.active_tasks,
worker.preparing_tasks,
worker.waiting_upstream_tasks,
worker.finalizing_tasks,
worker.pressure_state,
worker.pressure_reason,
worker.load_sampled_at,
count(DISTINCT task.id) FILTER (WHERE task.status = 'running')::int,
count(DISTINCT lease.id) FILTER (WHERE lease.released_at IS NULL)::int,
worker.heartbeat_at,
@@ -339,6 +433,16 @@ ORDER BY worker.site ASC, worker.status DESC, worker.instance_id ASC`,
&instance.Status,
&instance.Allocated,
&instance.CapacityLimit,
&instance.HardCapacityLimit,
&instance.SafeCapacity,
&instance.HeavyCapacity,
&instance.ReportedActiveTasks,
&instance.PreparingTasks,
&instance.WaitingUpstreamTasks,
&instance.FinalizingTasks,
&instance.PressureState,
&instance.PressureReason,
&instance.LoadSampledAt,
&instance.RunningTasks,
&instance.ActiveLeases,
&instance.HeartbeatAt,
@@ -476,7 +580,7 @@ WITH orphaned AS MATERIALIZED (
),
released_leases AS (
UPDATE gateway_concurrency_leases lease
SET released_at = now()
SET released_at = statement_timestamp()
FROM orphaned
WHERE lease.task_id = orphaned.task_id
AND lease.released_at IS NULL
@@ -42,3 +42,15 @@ func TestAllocateWorkerCapacitiesSupportsUnequalLimits(t *testing.T) {
t.Fatalf("allocations=%v global=%d, want 1/4 and 5", allocations, global)
}
}
func TestAllocateWorkerCapacitiesRedistributesFromPressuredWorker(t *testing.T) {
workers := []activeWorkerCapacity{
{InstanceID: "worker-a", CapacityLimit: 0},
{InstanceID: "worker-b", CapacityLimit: 6},
{InstanceID: "worker-c", CapacityLimit: 6},
}
allocations, global := allocateWorkerCapacities(workers, 8)
if global != 8 || allocations["worker-a"] != 0 || allocations["worker-b"] != 4 || allocations["worker-c"] != 4 {
t.Fatalf("allocations=%v global=%d, want 0/4/4 and 8", allocations, global)
}
}
+381
View File
@@ -0,0 +1,381 @@
package workerload
import (
"context"
"errors"
"strings"
"sync"
"time"
)
type Phase string
const (
PhasePreparing Phase = "preparing"
PhaseWaitingUpstream Phase = "waiting_upstream"
PhaseFinalizing Phase = "finalizing"
)
type PressureState string
const (
PressureNormal PressureState = "normal"
PressureBusy PressureState = "busy"
PressureCritical PressureState = "critical"
)
const (
ModeAdaptive = "adaptive"
ModeLegacy = "legacy"
)
var ErrReleased = errors.New("worker load lease already released")
type Config struct {
Mode string
HardLimit int
InitialActive int
InitialHeavy int
HealthySamples int
}
type ResourceSample struct {
MemoryCurrentBytes int64
MemoryLimitBytes int64
CPUUtilization float64
CPUThrottled bool
DBConnections int32
DBMaxConnections int32
SampledAt time.Time
}
type Snapshot struct {
Mode string
ActiveLimit int
HeavyLimit int
ClaimLimit int
SafeCapacity int
ActiveTasks int
PreparingTasks int
WaitingUpstreamTasks int
FinalizingTasks int
PressureState PressureState
PressureReason string
MemoryUtilization float64
CPUUtilization float64
DBUtilization float64
SampledAt time.Time
}
type Controller struct {
mu sync.Mutex
mode string
hardLimit int
activeLimit int
heavyLimit int
claimLimit int
healthySamples int
healthyCount int
preparing int
waiting int
finalizing int
last Snapshot
wake chan struct{}
}
type Lease struct {
controller *Controller
phase Phase
released bool
}
func New(config Config) *Controller {
mode := strings.ToLower(strings.TrimSpace(config.Mode))
if mode != ModeLegacy {
mode = ModeAdaptive
}
if config.HardLimit < 1 {
config.HardLimit = 1
}
if config.InitialActive < 1 {
config.InitialActive = 4
}
if config.InitialHeavy < 1 {
config.InitialHeavy = 1
}
if config.HealthySamples < 1 {
config.HealthySamples = 3
}
active := min(config.InitialActive, config.HardLimit)
heavy := min(config.InitialHeavy, active)
if mode == ModeLegacy {
active = config.HardLimit
heavy = config.HardLimit
}
controller := &Controller{
mode: mode, hardLimit: config.HardLimit,
activeLimit: active, heavyLimit: heavy, claimLimit: active,
healthySamples: config.HealthySamples,
wake: make(chan struct{}, 1),
}
controller.last = controller.snapshotLocked(ResourceSample{SampledAt: time.Now()})
return controller
}
func (c *Controller) Observe(sample ResourceSample) Snapshot {
c.mu.Lock()
defer c.mu.Unlock()
if sample.SampledAt.IsZero() {
sample.SampledAt = time.Now()
}
memory := utilization(sample.MemoryCurrentBytes, sample.MemoryLimitBytes)
database := utilization(int64(sample.DBConnections), int64(sample.DBMaxConnections))
cpu := bounded(sample.CPUUtilization)
state, reason := pressure(memory, cpu, database, sample.CPUThrottled)
if c.mode == ModeLegacy {
c.activeLimit = c.hardLimit
c.heavyLimit = c.hardLimit
state = PressureNormal
reason = "legacy"
} else {
c.adjustLocked(state, memory, cpu, database)
}
c.last = c.snapshotLocked(sample)
c.last.PressureState = state
c.last.PressureReason = reason
c.last.MemoryUtilization = memory
c.last.CPUUtilization = cpu
c.last.DBUtilization = database
c.last.SafeCapacity = c.activeLimit
if state == PressureCritical {
c.last.SafeCapacity = 0
}
c.signalLocked()
return c.last
}
func (c *Controller) SetClaimLimit(limit int) Snapshot {
c.mu.Lock()
defer c.mu.Unlock()
if limit < 0 {
limit = 0
}
c.claimLimit = min(limit, c.hardLimit)
c.last = c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
c.signalLocked()
return c.last
}
func (c *Controller) Snapshot() Snapshot {
c.mu.Lock()
defer c.mu.Unlock()
return c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
}
func (c *Controller) TryStart() (*Lease, bool) {
c.mu.Lock()
defer c.mu.Unlock()
activeLimit := min(c.activeLimit, c.claimLimit)
if activeLimit <= 0 || c.activeLocked() >= activeLimit || c.preparing+c.finalizing >= c.heavyLimit {
return nil, false
}
c.preparing++
c.last = c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
return &Lease{controller: c, phase: PhasePreparing}, true
}
func (l *Lease) EnterWaiting() error {
if l == nil || l.controller == nil {
return nil
}
c := l.controller
c.mu.Lock()
defer c.mu.Unlock()
if l.released {
return ErrReleased
}
if l.phase == PhaseWaitingUpstream {
return nil
}
c.decrementPhaseLocked(l.phase)
c.waiting++
l.phase = PhaseWaitingUpstream
c.last = c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
c.signalLocked()
return nil
}
func (l *Lease) EnterFinalizing(ctx context.Context) error {
if l == nil || l.controller == nil {
return nil
}
c := l.controller
for {
c.mu.Lock()
if l.released {
c.mu.Unlock()
return ErrReleased
}
if l.phase == PhaseFinalizing {
c.mu.Unlock()
return nil
}
// Even under critical pressure, let one submitted task at a time finish
// and release its provider lease instead of deadlocking the drain path.
heavyLimit := max(c.heavyLimit, 1)
if c.preparing+c.finalizing < heavyLimit {
c.decrementPhaseLocked(l.phase)
c.finalizing++
l.phase = PhaseFinalizing
c.last = c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
c.mu.Unlock()
return nil
}
wake := c.wake
c.mu.Unlock()
select {
case <-ctx.Done():
return ctx.Err()
case <-wake:
}
}
}
func (l *Lease) Release() {
if l == nil || l.controller == nil {
return
}
c := l.controller
c.mu.Lock()
defer c.mu.Unlock()
if l.released {
return
}
c.decrementPhaseLocked(l.phase)
l.released = true
c.last = c.snapshotLocked(ResourceSample{SampledAt: c.last.SampledAt})
c.signalLocked()
}
func (l *Lease) Phase() Phase {
if l == nil || l.controller == nil {
return ""
}
l.controller.mu.Lock()
defer l.controller.mu.Unlock()
return l.phase
}
func (c *Controller) adjustLocked(state PressureState, memory, cpu, database float64) {
switch state {
case PressureCritical:
c.healthyCount = 0
c.activeLimit = max(1, c.activeLimit/2)
c.heavyLimit = 1
case PressureBusy:
c.healthyCount = 0
step := max(1, c.activeLimit/4)
c.activeLimit = max(1, c.activeLimit-step)
c.heavyLimit = min(c.heavyLimit, max(1, (c.activeLimit+3)/4))
default:
if memory > .60 || cpu > .65 || database > .65 || c.activeLocked() < min(c.activeLimit, c.claimLimit) {
c.healthyCount = 0
return
}
c.healthyCount++
if c.healthyCount < c.healthySamples {
return
}
c.healthyCount = 0
step := max(1, c.activeLimit/4)
c.activeLimit = min(c.hardLimit, c.activeLimit+step)
c.heavyLimit = min(c.activeLimit, max(1, (c.activeLimit+3)/4))
}
}
func (c *Controller) snapshotLocked(sample ResourceSample) Snapshot {
sampledAt := sample.SampledAt
if sampledAt.IsZero() {
sampledAt = c.last.SampledAt
}
safeCapacity := c.activeLimit
if c.last.PressureState == PressureCritical {
safeCapacity = 0
}
return Snapshot{
Mode: c.mode,
ActiveLimit: c.activeLimit,
HeavyLimit: c.heavyLimit,
ClaimLimit: c.claimLimit,
SafeCapacity: safeCapacity,
ActiveTasks: c.activeLocked(),
PreparingTasks: c.preparing,
WaitingUpstreamTasks: c.waiting,
FinalizingTasks: c.finalizing,
PressureState: c.last.PressureState,
PressureReason: c.last.PressureReason,
MemoryUtilization: c.last.MemoryUtilization,
CPUUtilization: c.last.CPUUtilization,
DBUtilization: c.last.DBUtilization,
SampledAt: sampledAt,
}
}
func (c *Controller) activeLocked() int { return c.preparing + c.waiting + c.finalizing }
func (c *Controller) decrementPhaseLocked(phase Phase) {
switch phase {
case PhasePreparing:
c.preparing = max(0, c.preparing-1)
case PhaseWaitingUpstream:
c.waiting = max(0, c.waiting-1)
case PhaseFinalizing:
c.finalizing = max(0, c.finalizing-1)
}
}
func (c *Controller) signalLocked() {
select {
case c.wake <- struct{}{}:
default:
}
}
func pressure(memory, cpu, database float64, throttled bool) (PressureState, string) {
switch {
case memory >= .90:
return PressureCritical, "memory"
case database >= .90:
return PressureCritical, "database"
case cpu >= .95 && throttled:
return PressureCritical, "cpu_throttled"
case memory >= .75:
return PressureBusy, "memory"
case database >= .80:
return PressureBusy, "database"
case cpu >= .80 || throttled:
return PressureBusy, "cpu"
default:
return PressureNormal, ""
}
}
func utilization(current, limit int64) float64 {
if current <= 0 || limit <= 0 {
return 0
}
return bounded(float64(current) / float64(limit))
}
func bounded(value float64) float64 {
if value < 0 {
return 0
}
if value > 1 {
return 1
}
return value
}
@@ -0,0 +1,95 @@
package workerload
import (
"context"
"testing"
"time"
)
func TestControllerStartsConservativelyAndGrowsUnderSustainedDemand(t *testing.T) {
controller := New(Config{Mode: ModeAdaptive, HardLimit: 16, InitialActive: 4, InitialHeavy: 1, HealthySamples: 2})
leases := make([]*Lease, 0, 4)
for range 4 {
lease, ok := controller.TryStart()
if !ok {
t.Fatal("initial task was not admitted")
}
_ = lease.EnterWaiting()
leases = append(leases, lease)
}
for range 2 {
controller.Observe(ResourceSample{MemoryCurrentBytes: 40, MemoryLimitBytes: 100, CPUUtilization: .4, DBConnections: 4, DBMaxConnections: 20})
}
snapshot := controller.Snapshot()
if snapshot.ActiveLimit != 5 || snapshot.HeavyLimit != 2 {
t.Fatalf("grown snapshot=%+v, want active=5 heavy=2", snapshot)
}
for _, lease := range leases {
lease.Release()
}
}
func TestControllerBusyAndCriticalPressureReduceNewClaims(t *testing.T) {
controller := New(Config{Mode: ModeAdaptive, HardLimit: 16, InitialActive: 8, InitialHeavy: 2})
busy := controller.Observe(ResourceSample{MemoryCurrentBytes: 80, MemoryLimitBytes: 100})
if busy.PressureState != PressureBusy || busy.SafeCapacity >= 8 {
t.Fatalf("busy snapshot=%+v", busy)
}
critical := controller.Observe(ResourceSample{MemoryCurrentBytes: 95, MemoryLimitBytes: 100})
if critical.PressureState != PressureCritical || critical.SafeCapacity != 0 {
t.Fatalf("critical snapshot=%+v", critical)
}
controller.SetClaimLimit(0)
if _, ok := controller.TryStart(); ok {
t.Fatal("critical controller admitted a new task")
}
}
func TestWaitingReleasesHeavyPermitAndFinalizingReacquiresIt(t *testing.T) {
controller := New(Config{Mode: ModeAdaptive, HardLimit: 4, InitialActive: 4, InitialHeavy: 1})
first, ok := controller.TryStart()
if !ok {
t.Fatal("first lease unavailable")
}
if _, ok := controller.TryStart(); ok {
t.Fatal("second preparing task bypassed heavy limit")
}
if err := first.EnterWaiting(); err != nil {
t.Fatal(err)
}
second, ok := controller.TryStart()
if !ok {
t.Fatal("waiting task did not release heavy permit")
}
if err := second.EnterWaiting(); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
if err := first.EnterFinalizing(ctx); err != nil {
t.Fatal(err)
}
blocked := make(chan error, 1)
go func() { blocked <- second.EnterFinalizing(ctx) }()
select {
case err := <-blocked:
t.Fatalf("second finalizer did not wait: %v", err)
case <-time.After(20 * time.Millisecond):
}
first.Release()
if err := <-blocked; err != nil {
t.Fatal(err)
}
second.Release()
if snapshot := controller.Snapshot(); snapshot.ActiveTasks != 0 {
t.Fatalf("active tasks=%d, want 0", snapshot.ActiveTasks)
}
}
func TestLegacyModeUsesHardLimit(t *testing.T) {
controller := New(Config{Mode: ModeLegacy, HardLimit: 7})
snapshot := controller.Observe(ResourceSample{MemoryCurrentBytes: 99, MemoryLimitBytes: 100, CPUUtilization: 1, CPUThrottled: true})
if snapshot.ActiveLimit != 7 || snapshot.HeavyLimit != 7 || snapshot.SafeCapacity != 7 || snapshot.PressureReason != "legacy" {
t.Fatalf("legacy snapshot=%+v", snapshot)
}
}
+131
View File
@@ -0,0 +1,131 @@
package workerload
import (
"errors"
"os"
"runtime"
"strconv"
"strings"
"sync"
"time"
)
type SystemSampler struct {
mu sync.Mutex
root string
lastCPUUsage int64
lastCPUTime time.Time
lastThrottled int64
}
func NewSystemSampler() *SystemSampler {
return &SystemSampler{root: "/sys/fs/cgroup"}
}
func NewSystemSamplerAt(root string) *SystemSampler {
return &SystemSampler{root: strings.TrimRight(root, "/")}
}
func (s *SystemSampler) Sample(databaseConnections, databaseMax int32) ResourceSample {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
memoryCurrent, memoryLimit := s.memory()
cpuUsage, throttledCount := s.cpuStat()
cpuLimit := s.cpuLimit()
cpuUtilization := 0.0
if s.lastCPUUsage > 0 && cpuUsage >= s.lastCPUUsage && !s.lastCPUTime.IsZero() {
elapsed := now.Sub(s.lastCPUTime).Seconds()
if elapsed > 0 {
cpuUtilization = float64(cpuUsage-s.lastCPUUsage) / 1_000_000 / elapsed / cpuLimit
}
}
throttled := s.lastThrottled > 0 && throttledCount > s.lastThrottled
s.lastCPUUsage = cpuUsage
s.lastThrottled = throttledCount
s.lastCPUTime = now
return ResourceSample{
MemoryCurrentBytes: memoryCurrent,
MemoryLimitBytes: memoryLimit,
CPUUtilization: bounded(cpuUtilization),
CPUThrottled: throttled,
DBConnections: databaseConnections,
DBMaxConnections: databaseMax,
SampledAt: now,
}
}
func (s *SystemSampler) memory() (int64, int64) {
current, currentErr := readIntFile(s.root + "/memory.current")
limit, limitErr := readLimitFile(s.root + "/memory.max")
if currentErr == nil && limitErr == nil {
return current, limit
}
current, _ = readIntFile(s.root + "/memory/memory.usage_in_bytes")
limit, _ = readLimitFile(s.root + "/memory/memory.limit_in_bytes")
return current, limit
}
func (s *SystemSampler) cpuStat() (usageUsec, throttled int64) {
data, err := os.ReadFile(s.root + "/cpu.stat")
if err != nil {
return 0, 0
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) != 2 {
continue
}
value, parseErr := strconv.ParseInt(fields[1], 10, 64)
if parseErr != nil {
continue
}
switch fields[0] {
case "usage_usec":
usageUsec = value
case "nr_throttled":
throttled = value
}
}
return usageUsec, throttled
}
func (s *SystemSampler) cpuLimit() float64 {
data, err := os.ReadFile(s.root + "/cpu.max")
if err == nil {
fields := strings.Fields(string(data))
if len(fields) == 2 && fields[0] != "max" {
quota, quotaErr := strconv.ParseFloat(fields[0], 64)
period, periodErr := strconv.ParseFloat(fields[1], 64)
if quotaErr == nil && periodErr == nil && quota > 0 && period > 0 {
return max(quota/period, .001)
}
}
}
return max(float64(runtime.GOMAXPROCS(0)), 1)
}
func readIntFile(path string) (int64, error) {
data, err := os.ReadFile(path)
if err != nil {
return 0, err
}
return strconv.ParseInt(strings.TrimSpace(string(data)), 10, 64)
}
func readLimitFile(path string) (int64, error) {
data, err := os.ReadFile(path)
if err != nil {
return 0, err
}
value := strings.TrimSpace(string(data))
if value == "" || value == "max" {
return 0, errors.New("cgroup limit is unlimited")
}
limit, err := strconv.ParseInt(value, 10, 64)
if err != nil || limit <= 0 || limit > 1<<60 {
return 0, errors.New("cgroup limit is not finite")
}
return limit, nil
}
@@ -0,0 +1,34 @@
package workerload
import (
"os"
"path/filepath"
"testing"
"time"
)
func TestSystemSamplerReadsCgroupV2(t *testing.T) {
root := t.TempDir()
writeSamplerFixture(t, filepath.Join(root, "memory.current"), "50\n")
writeSamplerFixture(t, filepath.Join(root, "memory.max"), "100\n")
writeSamplerFixture(t, filepath.Join(root, "cpu.max"), "100000 100000\n")
writeSamplerFixture(t, filepath.Join(root, "cpu.stat"), "usage_usec 100000\nnr_throttled 1\n")
sampler := NewSystemSamplerAt(root)
first := sampler.Sample(2, 10)
if first.MemoryCurrentBytes != 50 || first.MemoryLimitBytes != 100 || first.DBConnections != 2 {
t.Fatalf("first sample=%+v", first)
}
time.Sleep(10 * time.Millisecond)
writeSamplerFixture(t, filepath.Join(root, "cpu.stat"), "usage_usec 105000\nnr_throttled 2\n")
second := sampler.Sample(3, 10)
if second.CPUUtilization <= 0 || !second.CPUThrottled {
t.Fatalf("second sample=%+v", second)
}
}
func writeSamplerFixture(t *testing.T, path, value string) {
t.Helper()
if err := os.WriteFile(path, []byte(value), 0o600); err != nil {
t.Fatal(err)
}
}
@@ -0,0 +1,44 @@
ALTER TABLE gateway_worker_instances
ADD COLUMN hard_capacity_limit integer NOT NULL DEFAULT 0,
ADD COLUMN safe_capacity integer NOT NULL DEFAULT 0,
ADD COLUMN heavy_capacity integer NOT NULL DEFAULT 0,
ADD COLUMN active_tasks integer NOT NULL DEFAULT 0,
ADD COLUMN preparing_tasks integer NOT NULL DEFAULT 0,
ADD COLUMN waiting_upstream_tasks integer NOT NULL DEFAULT 0,
ADD COLUMN finalizing_tasks integer NOT NULL DEFAULT 0,
ADD COLUMN pressure_state text NOT NULL DEFAULT 'unknown',
ADD COLUMN pressure_reason text NOT NULL DEFAULT '',
ADD COLUMN load_sampled_at timestamptz,
ADD CONSTRAINT gateway_worker_instances_adaptive_capacity_check
CHECK (
hard_capacity_limit >= 0
AND safe_capacity >= 0
AND heavy_capacity >= 0
AND active_tasks >= 0
AND preparing_tasks >= 0
AND waiting_upstream_tasks >= 0
AND finalizing_tasks >= 0
AND active_tasks = preparing_tasks + waiting_upstream_tasks + finalizing_tasks
) NOT VALID,
ADD CONSTRAINT gateway_worker_instances_pressure_state_check
CHECK (pressure_state IN ('unknown', 'normal', 'busy', 'critical')) NOT VALID;
ALTER TABLE gateway_concurrency_leases
ADD COLUMN limit_value numeric,
ADD CONSTRAINT gateway_concurrency_leases_limit_value_check
CHECK (limit_value IS NULL OR limit_value > 0) NOT VALID;
UPDATE gateway_worker_instances
SET hard_capacity_limit = capacity_limit,
safe_capacity = capacity_limit,
heavy_capacity = capacity_limit
WHERE hard_capacity_limit = 0;
ALTER TABLE gateway_worker_instances
VALIDATE CONSTRAINT gateway_worker_instances_adaptive_capacity_check;
ALTER TABLE gateway_worker_instances
VALIDATE CONSTRAINT gateway_worker_instances_pressure_state_check;
ALTER TABLE gateway_concurrency_leases
VALIDATE CONSTRAINT gateway_concurrency_leases_limit_value_check;
+16 -4
View File
@@ -37,6 +37,7 @@ import type {
UserGroupUpsertRequest,
UserGroup,
WalletRechargeRequest,
WorkerClusterRuntime,
} from '@easyai-ai-gateway/contracts';
import {
batchAccessRules,
@@ -67,6 +68,7 @@ import {
getRunnerPolicy,
getSecurityEventConnection,
getWalletSummary,
getWorkerClusterRuntime,
listAccessRules,
listAdminTasks,
listAuditLogs,
@@ -201,6 +203,7 @@ type DataKey =
| 'runtimePolicySets'
| 'rateLimitWindows'
| 'modelRateLimits'
| 'workerClusterRuntime'
| 'tenants'
| 'users'
| 'userGroups'
@@ -255,6 +258,7 @@ export function App() {
const [rateLimitWindows, setRateLimitWindows] = useState<RateLimitWindow[]>([]);
const [modelRateLimits, setModelRateLimits] = useState<ModelRateLimitStatus[]>([]);
const [modelRateLimitsUpdatedAt, setModelRateLimitsUpdatedAt] = useState<number | null>(null);
const [workerClusterRuntime, setWorkerClusterRuntime] = useState<WorkerClusterRuntime | null>(null);
const [tenants, setTenants] = useState<GatewayTenant[]>([]);
const [users, setUsers] = useState<GatewayUser[]>([]);
const [userGroups, setUserGroups] = useState<UserGroup[]>([]);
@@ -381,18 +385,21 @@ export function App() {
useEffect(() => {
if (!token || activePage !== 'admin' || adminSection !== 'realtimeLoad') return undefined;
const timer = window.setInterval(() => {
void Promise.all([listModelRateLimitStatuses(token), listPlatforms(token)])
.then(([rateLimitResponse, platformResponse]) => {
void Promise.all([listModelRateLimitStatuses(token), listPlatforms(token), getWorkerClusterRuntime(token)])
.then(([rateLimitResponse, platformResponse, workerRuntime]) => {
setModelRateLimits(rateLimitResponse.items);
setModelRateLimitsUpdatedAt(Date.now());
setPlatforms(platformResponse.items);
setWorkerClusterRuntime(workerRuntime);
loadedDataKeysRef.current.add('modelRateLimits');
loadedDataKeysRef.current.add('platforms');
loadedDataKeysRef.current.add('workerClusterRuntime');
})
.catch((err) => {
if (handleAuthExpired(err, token)) return;
loadedDataKeysRef.current.delete('modelRateLimits');
loadedDataKeysRef.current.delete('platforms');
loadedDataKeysRef.current.delete('workerClusterRuntime');
});
}, 3000);
return () => window.clearInterval(timer);
@@ -446,6 +453,7 @@ export function App() {
rateLimitWindows,
modelRateLimits,
modelRateLimitsUpdatedAt,
workerClusterRuntime,
runtimePolicySets,
securityEventConnection,
taskResult,
@@ -455,7 +463,7 @@ export function App() {
users,
walletAccounts,
walletTransactions,
}), [accessRules, adminTasks, apiKeys, auditLogs, baseModels, clientCustomizationSettings, currentUser, currentUserGroups, fileStorageChannels, fileStorageSettings, modelCatalog, modelRateLimits, modelRateLimitsUpdatedAt, models, networkProxyConfig, platforms, pricingRuleSets, pricingRules, providers, rateLimitWindows, runnerPolicy, runtimePolicySets, securityEventConnection, taskResult, tasks, tenants, userGroups, users, walletAccounts, walletTransactions]);
}), [accessRules, adminTasks, apiKeys, auditLogs, baseModels, clientCustomizationSettings, currentUser, currentUserGroups, fileStorageChannels, fileStorageSettings, modelCatalog, modelRateLimits, modelRateLimitsUpdatedAt, models, networkProxyConfig, platforms, pricingRuleSets, pricingRules, providers, rateLimitWindows, runnerPolicy, runtimePolicySets, securityEventConnection, taskResult, tasks, tenants, userGroups, users, walletAccounts, walletTransactions, workerClusterRuntime]);
async function refresh(nextToken = token) {
await ensureRouteData(nextToken, true);
@@ -593,6 +601,9 @@ export function App() {
setModelRateLimitsUpdatedAt(Date.now());
}
return;
case 'workerClusterRuntime':
setWorkerClusterRuntime(await getWorkerClusterRuntime(nextToken));
return;
case 'tenants':
setTenants((await listTenants(nextToken)).items);
return;
@@ -1255,6 +1266,7 @@ export function App() {
setAuditLogs([]);
setRateLimitWindows([]);
setModelRateLimits([]);
setWorkerClusterRuntime(null);
setTenants([]);
setUsers([]);
setUserGroups([]);
@@ -1744,7 +1756,7 @@ function dataKeysForRoute(
case 'platforms':
return ['platforms', 'models', 'providers', 'baseModels', 'pricingRuleSets', 'networkProxyConfig'];
case 'realtimeLoad':
return ['platforms', 'modelRateLimits'];
return ['platforms', 'modelRateLimits', 'workerClusterRuntime'];
case 'tasks':
return ['adminTasks', 'tenants', 'users', 'userGroups', 'platforms', 'models'];
case 'tenants':
+5
View File
@@ -59,6 +59,7 @@ import type {
WalletBalanceAdjustmentRequest,
WalletRechargeRequest,
WalletSummaryResponse,
WorkerClusterRuntime,
} from '@easyai-ai-gateway/contracts';
import type { AdminTaskQuery, PlatformCreateInput, PlatformModelBindingInput, WorkspaceTaskQuery } from './types';
@@ -1037,6 +1038,10 @@ export async function listModelRateLimitStatuses(token: string): Promise<ListRes
return request<ListResponse<ModelRateLimitStatus>>('/api/admin/runtime/model-rate-limits', { token });
}
export async function getWorkerClusterRuntime(token: string): Promise<WorkerClusterRuntime> {
return request<WorkerClusterRuntime>('/api/admin/runtime/workers', { token });
}
export async function restoreModelRuntimeStatus(token: string, platformModelId: string): Promise<ModelRateLimitStatus> {
return request<ModelRateLimitStatus>(`/api/admin/runtime/model-rate-limits/${platformModelId}/restore`, {
method: 'POST',
+2
View File
@@ -26,6 +26,7 @@ import type {
RuntimePolicySet,
SecurityEventConnectionResponse,
UserGroup,
WorkerClusterRuntime,
} from '@easyai-ai-gateway/contracts';
export interface ConsoleData {
@@ -50,6 +51,7 @@ export interface ConsoleData {
rateLimitWindows: RateLimitWindow[];
modelRateLimits: ModelRateLimitStatus[];
modelRateLimitsUpdatedAt: number | null;
workerClusterRuntime: WorkerClusterRuntime | null;
runtimePolicySets: RuntimePolicySet[];
securityEventConnection: SecurityEventConnectionResponse | null;
taskResult: GatewayTask | null;
+1 -1
View File
@@ -50,7 +50,7 @@ export const adminPages = [
{ title: '用户组策略', path: '/admin/user-groups', description: '用户组成员、充值折扣、调用折扣、TPM/RPM/并发和队列优先级。' },
{ title: '全局模型配置', path: '/admin/models/global', description: '基准模型库、能力 schema、基准定价和默认限流模板。' },
{ title: '平台管理', path: '/admin/platforms', description: '平台 CRUD、凭证、默认折扣、平台模型、限流和重试策略。' },
{ title: '实时负载', path: '/admin/realtime-load', description: '平台模型查看实时 RPM、TPM、并发、排队和冷却状态。' },
{ title: '实时负载', path: '/admin/realtime-load', description: '查看平台模型 RPM、TPM、并发,以及 Worker 自适应容量和压力状态。' },
{ title: '任务记录', path: '/admin/tasks', description: '跨租户查询任务、执行链路、参数转换、计费和原始详情。' },
{ title: '计费结算', path: '/admin/billing-settlements', description: '查询计费结算队列,批量处理等待重试和人工复核记录。' },
{ title: '运行与队列', path: '/admin/runtime/queues', description: 'TPM/RPM 窗口、并发 lease、cooldown、任务恢复和队列积压。' },
+1
View File
@@ -179,6 +179,7 @@ export function AdminPage(props: {
modelRateLimits={props.data.modelRateLimits}
modelRateLimitsUpdatedAt={props.data.modelRateLimitsUpdatedAt}
platforms={props.data.platforms}
workerClusterRuntime={props.data.workerClusterRuntime}
onSavePlatformDynamicPriority={props.onSavePlatformDynamicPriority}
onRestoreRuntimeModel={props.onRestoreRuntimeModel}
/>
+92 -1
View File
@@ -1,13 +1,14 @@
import { useEffect, useMemo, useState, type FormEvent } from 'react';
import { Popover as AntPopover } from 'antd';
import { CheckCircle2, History, RotateCcw, Search, SlidersHorizontal } from 'lucide-react';
import type { IntegrationPlatform, ModelRateLimitStatus, PlatformDynamicPriorityUpdateRequest, PlatformPolicyEvent, PriorityDemotionRecord } from '@easyai-ai-gateway/contracts';
import type { IntegrationPlatform, ModelRateLimitStatus, PlatformDynamicPriorityUpdateRequest, PlatformPolicyEvent, PriorityDemotionRecord, WorkerClusterRuntime, WorkerInstanceRuntime } from '@easyai-ai-gateway/contracts';
import { Badge, Button, Card, CardContent, EmptyState, FormDialog, Input, Label, Select, Table, TableCell, TableHead, TableRow } from '../../components/ui';
export function RealtimeLoadPanel(props: {
modelRateLimits: ModelRateLimitStatus[];
modelRateLimitsUpdatedAt: number | null;
platforms: IntegrationPlatform[];
workerClusterRuntime: WorkerClusterRuntime | null;
onSavePlatformDynamicPriority: (platformId: string, input: PlatformDynamicPriorityUpdateRequest) => Promise<void>;
onRestoreRuntimeModel: (platformModelId: string) => Promise<void>;
}) {
@@ -119,6 +120,7 @@ export function RealtimeLoadPanel(props: {
return (
<section className="pageStack">
<WorkerRuntimeTable runtime={props.workerClusterRuntime} />
<Card className="compactAdminTableCard">
<CardContent className="compactAdminTableContent">
<div className="compactAdminToolbar realtimeCompactToolbar">
@@ -241,6 +243,95 @@ export function RealtimeLoadPanel(props: {
);
}
function WorkerRuntimeTable(props: { runtime: WorkerClusterRuntime | null }) {
const workers = props.runtime?.workers ?? [];
const queue = props.runtime?.queue;
return (
<Card className="compactAdminTableCard">
<CardContent className="compactAdminTableContent">
<div className="compactAdminToolbar">
<span className="platformTableName">
<strong>Worker </strong>
<small>
{queue
? `共享队列 ${queue.queued},运行 ${queue.running},最老等待 ${Math.round(queue.oldestWaitSeconds)}`
: '等待集群负载快照'}
</small>
</span>
</div>
{!workers.length ? (
<EmptyState title="暂无活跃 Worker" description="Worker 心跳后会显示安全容量、阶段分布和压力状态。" />
) : (
<div className="platformLimitTableViewport">
<Table className="platformDataTable platformLimitTable" density="compact">
<TableRow className="shTableHeader">
<TableHead>Worker</TableHead>
<TableHead></TableHead>
<TableHead className="platformLimitNumberHead"> / / Hard</TableHead>
<TableHead className="platformLimitNumberHead"></TableHead>
<TableHead className="platformLimitNumberHead"> / / </TableHead>
<TableHead className="platformLimitNumberHead"> / </TableHead>
<TableHead></TableHead>
</TableRow>
{workers.map((worker) => (
<TableRow key={worker.instanceId}>
<TableCell>
<span className="platformTableName">
<strong>{worker.podName || worker.instanceId}</strong>
<small>{[worker.site, shortId(worker.revision)].filter(Boolean).join(' · ') || shortId(worker.instanceId)}</small>
</span>
</TableCell>
<TableCell>{workerPressureCell(worker)}</TableCell>
<TableCell className="platformLimitNumberCell">
<strong>{worker.allocatedCapacity} / {worker.safeCapacity} / {worker.hardCapacityLimit}</strong>
</TableCell>
<TableCell className="platformLimitNumberCell">{worker.heavyCapacity}</TableCell>
<TableCell className="platformLimitNumberCell">
<span className="rateMetricCell">
<strong>{worker.preparingTasks} / {worker.waitingUpstreamTasks} / {worker.finalizingTasks}</strong>
<small> {worker.reportedActiveTasks}</small>
</span>
</TableCell>
<TableCell className="platformLimitNumberCell">{worker.runningTasks} / {worker.activeLeases}</TableCell>
<TableCell>
<span className="platformTableName">
<strong>{formatDateTime(worker.loadSampledAt) || '-'}</strong>
<small> {formatDateTime(worker.heartbeatAt)}</small>
</span>
</TableCell>
</TableRow>
))}
</Table>
</div>
)}
</CardContent>
</Card>
);
}
function workerPressureCell(worker: WorkerInstanceRuntime) {
const variant = worker.pressureState === 'critical'
? 'destructive'
: worker.pressureState === 'busy'
? 'warning'
: worker.pressureState === 'normal'
? 'success'
: 'secondary';
const label = worker.pressureState === 'critical'
? '临界'
: worker.pressureState === 'busy'
? '繁忙'
: worker.pressureState === 'normal'
? '正常'
: '未知';
return (
<span className="platformTableName">
<strong><Badge variant={variant}>{label}</Badge></strong>
<small>{worker.pressureReason || '无压力原因'}</small>
</span>
);
}
type PriorityDialogState = {
platform: IntegrationPlatform | undefined;
status: ModelRateLimitStatus;