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:
@@ -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))
|
||||
|
||||
@@ -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" {
|
||||
|
||||
@@ -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": {
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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++ {
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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`,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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(¤t); 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 == "" {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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':
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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、任务恢复和队列积压。' },
|
||||
|
||||
@@ -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}
|
||||
/>
|
||||
|
||||
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user