Merge origin/main into codex/gemini-embedding-batch

# Conflicts:
#	apps/aether-gateway/src/ai_serving/api.rs
#	apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs
#	apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs
#	apps/aether-gateway/src/ai_serving/transport.rs
#	apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs
#	apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs
#	crates/aether-data/src/repository/candidate_selection/postgres.rs
#	crates/aether-model-fetch/src/strategy.rs
This commit is contained in:
MMEXA
2026-05-18 19:02:19 +00:00
446 changed files with 53043 additions and 4780 deletions

View File

@@ -58,6 +58,11 @@ ADMIN_USERNAME=admin
# AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true # AETHER_GATEWAY_AUTO_PREPARE_DATABASE=true
# PostgreSQL 连接池配置(默认适合单实例/小型部署;高并发可按需调大) # PostgreSQL 连接池配置(默认适合单实例/小型部署;高并发可按需调大)
# 推荐计算方式(单实例):
# MAX = CPU 核数 × 10AI 网关偏 IO 等待,可激进些;纯 OLTP 用 × 4
# MIN = MAX × 0.2(保留常驻连接应对突发流量,避免冷启动握手开销)
# 多实例部署时请按 实例数 × MAX 控制总和PG 端 max_connections 至少为该总和 + 20 余量
# AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=1 # AETHER_GATEWAY_DATA_POSTGRES_MIN_CONNECTIONS=1
# AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=20 # AETHER_GATEWAY_DATA_POSTGRES_MAX_CONNECTIONS=20
# AETHER_GATEWAY_DATA_POSTGRES_IDLE_TIMEOUT_MS=30000 # AETHER_GATEWAY_DATA_POSTGRES_STATEMENT_CACHE_CAPACITY=100
# AETHER_GATEWAY_DATA_POSTGRES_ACQUIRE_TIMEOUT_MS=3000

View File

@@ -247,8 +247,10 @@ jobs:
set -euo pipefail set -euo pipefail
if [[ "${GITHUB_REF_TYPE}" == "tag" ]]; then if [[ "${GITHUB_REF_TYPE}" == "tag" ]]; then
VERSION="${GITHUB_REF_NAME}" VERSION="${GITHUB_REF_NAME}"
SOURCE_REF="${GITHUB_REF_NAME}"
else else
VERSION="snapshot-${GITHUB_SHA::7}" VERSION="snapshot-${GITHUB_SHA::7}"
SOURCE_REF="${GITHUB_SHA}"
fi fi
mkdir -p package release-assets mkdir -p package release-assets
@@ -258,13 +260,20 @@ jobs:
root="package/${bundle}" root="package/${bundle}"
mkdir -p \ mkdir -p \
"${root}/bin" \ "${root}/bin" \
"${root}/frontend" "${root}/frontend" \
"${root}/scripts"
install -m 0755 "artifacts/aether-gateway-${platform}-${arch}/aether-gateway" "${root}/bin/aether-gateway" install -m 0755 "artifacts/aether-gateway-${platform}-${arch}/aether-gateway" "${root}/bin/aether-gateway"
cp -R artifacts/frontend-dist/. "${root}/frontend/" cp -R artifacts/frontend-dist/. "${root}/frontend/"
sed "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" install.sh > "${root}/install.sh" sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > "${root}/install.sh"
chmod 0755 "${root}/install.sh" chmod 0755 "${root}/install.sh"
install -m 0644 docker-compose.yml "${root}/docker-compose.yml" install -m 0644 docker-compose.yml "${root}/docker-compose.yml"
install -m 0644 docker-compose.single-node.yml "${root}/docker-compose.single-node.yml"
install -m 0755 scripts/migrate-pg-compose-to-single-node.sh "${root}/scripts/migrate-pg-compose-to-single-node.sh"
install -m 0755 scripts/migrate-pg-to-single-node.sh "${root}/scripts/migrate-pg-to-single-node.sh"
install -m 0644 .env.example "${root}/.env.example" install -m 0644 .env.example "${root}/.env.example"
install -m 0755 generate_keys.sh "${root}/generate_keys.sh" install -m 0755 generate_keys.sh "${root}/generate_keys.sh"
install -m 0644 README.md "${root}/README.md" install -m 0644 README.md "${root}/README.md"
@@ -274,7 +283,10 @@ jobs:
done done
done done
sed "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" install.sh > release-assets/install.sh sed \
-e "s/^SOURCE_REF=\"\${AETHER_SOURCE_REF:-main}\"/SOURCE_REF=\"\${AETHER_SOURCE_REF:-${SOURCE_REF}}\"/" \
-e "s/^VERSION=\"\${AETHER_VERSION:-}\"/VERSION=\"\${AETHER_VERSION:-${VERSION}}\"/" \
install.sh > release-assets/install.sh
chmod +x release-assets/install.sh chmod +x release-assets/install.sh
(cd release-assets && sha256sum *.tar.gz > SHA256SUMS) (cd release-assets && sha256sum *.tar.gz > SHA256SUMS)

380
Cargo.lock generated
View File

@@ -29,6 +29,7 @@ dependencies = [
"aether-data", "aether-data",
"aether-data-contracts", "aether-data-contracts",
"aether-provider-pool", "aether-provider-pool",
"aether-provider-transport",
"axum", "axum",
"base64 0.22.1", "base64 0.22.1",
"chrono", "chrono",
@@ -126,6 +127,7 @@ dependencies = [
"aether-ai-formats", "aether-ai-formats",
"aether-cache", "aether-cache",
"aether-data-contracts", "aether-data-contracts",
"aether-data-query",
"aether-wallet", "aether-wallet",
"async-trait", "async-trait",
"chrono", "chrono",
@@ -148,6 +150,7 @@ name = "aether-data-contracts"
version = "0.1.0" version = "0.1.0"
dependencies = [ dependencies = [
"aether-ai-formats", "aether-ai-formats",
"aether-routing-core",
"async-trait", "async-trait",
"chrono", "chrono",
"serde", "serde",
@@ -155,6 +158,13 @@ dependencies = [
"thiserror 2.0.18", "thiserror 2.0.18",
] ]
[[package]]
name = "aether-data-query"
version = "0.1.0"
dependencies = [
"sqlx",
]
[[package]] [[package]]
name = "aether-data-schema" name = "aether-data-schema"
version = "0.1.0" version = "0.1.0"
@@ -196,6 +206,7 @@ dependencies = [
"aether-pool-core", "aether-pool-core",
"aether-provider-pool", "aether-provider-pool",
"aether-provider-transport", "aether-provider-transport",
"aether-routing-core",
"aether-runtime", "aether-runtime",
"aether-runtime-state", "aether-runtime-state",
"aether-scheduler-core", "aether-scheduler-core",
@@ -240,6 +251,8 @@ dependencies = [
"url", "url",
"uuid", "uuid",
"webpki-roots 0.26.11", "webpki-roots 0.26.11",
"wreq",
"wreq-util",
] ]
[[package]] [[package]]
@@ -377,6 +390,16 @@ dependencies = [
"webpki-roots 0.26.11", "webpki-roots 0.26.11",
] ]
[[package]]
name = "aether-routing-core"
version = "0.1.0"
dependencies = [
"regex",
"serde",
"serde_json",
"thiserror 2.0.18",
]
[[package]] [[package]]
name = "aether-runtime" name = "aether-runtime"
version = "0.1.0" version = "0.1.0"
@@ -444,6 +467,7 @@ version = "0.1.0"
dependencies = [ dependencies = [
"aether-contracts", "aether-contracts",
"aether-data", "aether-data",
"aether-data-contracts",
"aether-gateway", "aether-gateway",
"aether-http", "aether-http",
"aether-runtime", "aether-runtime",
@@ -498,6 +522,18 @@ dependencies = [
"serde", "serde",
] ]
[[package]]
name = "ahash"
version = "0.8.12"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5a15f179cd60c4584b8a8c596927aadc462e27f2ca70c04e0071964a73ba7a75"
dependencies = [
"cfg-if",
"once_cell",
"version_check",
"zerocopy",
]
[[package]] [[package]]
name = "aho-corasick" name = "aho-corasick"
version = "1.1.4" version = "1.1.4"
@@ -507,6 +543,21 @@ dependencies = [
"memchr", "memchr",
] ]
[[package]]
name = "alloc-no-stdlib"
version = "2.0.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cc7bb162ec39d46ab1ca8c77bf72e890535becd1751bb45f64c597edb4c8c6b3"
[[package]]
name = "alloc-stdlib"
version = "0.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "94fb8275041c72129eb51b7d0322c29b8387a0386127718b096429201a5d6ece"
dependencies = [
"alloc-no-stdlib",
]
[[package]] [[package]]
name = "allocator-api2" name = "allocator-api2"
version = "0.2.21" version = "0.2.21"
@@ -809,6 +860,24 @@ dependencies = [
"zeroize", "zeroize",
] ]
[[package]]
name = "bindgen"
version = "0.72.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "993776b509cfb49c750f11b8f07a46fa23e0a1386ffc01fb1e7d343efc387895"
dependencies = [
"bitflags 2.11.0",
"cexpr",
"clang-sys",
"itertools 0.13.0",
"proc-macro2",
"quote",
"regex",
"rustc-hash",
"shlex",
"syn 2.0.117",
]
[[package]] [[package]]
name = "bit-set" name = "bit-set"
version = "0.5.3" version = "0.5.3"
@@ -867,6 +936,52 @@ dependencies = [
"cipher", "cipher",
] ]
[[package]]
name = "boring-sys2"
version = "5.0.0-alpha.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "455d79965f5155dcc88a7abce112c3590883889131b799beda10bf9a813ed669"
dependencies = [
"bindgen",
"cmake",
"fs_extra",
"fslock",
]
[[package]]
name = "boring2"
version = "5.0.0-alpha.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "183ccc3854411c035410dcdbffafca62084f3a6c33f013c77e83c025d2a08a28"
dependencies = [
"bitflags 2.11.0",
"boring-sys2",
"foreign-types",
"libc",
"openssl-macros",
]
[[package]]
name = "brotli"
version = "8.0.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4bd8b9603c7aa97359dbd97ecf258968c95f3adddd6db2f7e7a5bef101c84560"
dependencies = [
"alloc-no-stdlib",
"alloc-stdlib",
"brotli-decompressor",
]
[[package]]
name = "brotli-decompressor"
version = "5.0.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "874bb8112abecc98cbd6d81ea4fa7e94fb9449648c93cc89aa40c81c24d7de03"
dependencies = [
"alloc-no-stdlib",
"alloc-stdlib",
]
[[package]] [[package]]
name = "bumpalo" name = "bumpalo"
version = "3.20.2" version = "3.20.2"
@@ -921,6 +1036,15 @@ dependencies = [
"shlex", "shlex",
] ]
[[package]]
name = "cexpr"
version = "0.6.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6fac387a98bb7c37292057cffc56d62ecb629900026402633ae9160df93a8766"
dependencies = [
"nom",
]
[[package]] [[package]]
name = "cfg-if" name = "cfg-if"
version = "1.0.4" version = "1.0.4"
@@ -967,6 +1091,17 @@ dependencies = [
"inout", "inout",
] ]
[[package]]
name = "clang-sys"
version = "1.8.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0b023947811758c97c59bf9d1c188fd619ad4718dcaa767947df1cadb14f39f4"
dependencies = [
"glob",
"libc",
"libloading",
]
[[package]] [[package]]
name = "clap" name = "clap"
version = "4.6.0" version = "4.6.0"
@@ -1542,6 +1677,33 @@ version = "0.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb" checksum = "77ce24cb58228fbb8aa041425bb1050850ac19177686ea6e0f41a70416f56fdb"
[[package]]
name = "foreign-types"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d737d9aa519fb7b749cbc3b962edcf310a8dd1f4b67c91c4f83975dbdd17d965"
dependencies = [
"foreign-types-macros",
"foreign-types-shared",
]
[[package]]
name = "foreign-types-macros"
version = "0.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1a5c6c585bc94aaf2c7b51dd4c2ba22680844aba4c687be581871a6f518c5742"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]]
name = "foreign-types-shared"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "aa9a19cbb55df58761df49b23516a86d432839add4af60fc256da840f66ed35b"
[[package]] [[package]]
name = "form_urlencoded" name = "form_urlencoded"
version = "1.2.2" version = "1.2.2"
@@ -1557,6 +1719,16 @@ version = "1.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c"
[[package]]
name = "fslock"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "04412b8935272e3a9bae6f48c7bfff74c2911f60525404edfdd28e49884c3bfb"
dependencies = [
"libc",
"winapi",
]
[[package]] [[package]]
name = "futures" name = "futures"
version = "0.3.32" version = "0.3.32"
@@ -1706,6 +1878,12 @@ dependencies = [
"wasip3", "wasip3",
] ]
[[package]]
name = "glob"
version = "0.3.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280"
[[package]] [[package]]
name = "h2" name = "h2"
version = "0.4.13" version = "0.4.13"
@@ -1725,6 +1903,12 @@ dependencies = [
"tracing", "tracing",
] ]
[[package]]
name = "hashbrown"
version = "0.13.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "43a3c133739dddd0d2990f9a4bdf8eb4b21ef50e4851ca85ab661199821d510e"
[[package]] [[package]]
name = "hashbrown" name = "hashbrown"
version = "0.14.5" version = "0.14.5"
@@ -1840,6 +2024,26 @@ version = "0.4.2"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9171a2ea8a68358193d15dd5d70c1c10a2afc3e7e4c5bc92bc9f025cebd7359c" checksum = "9171a2ea8a68358193d15dd5d70c1c10a2afc3e7e4c5bc92bc9f025cebd7359c"
[[package]]
name = "http2"
version = "0.5.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "569ef7a780e853c4e1768f58a3c8168193b82cdcbab66638a0b1c6583ec5995e"
dependencies = [
"atomic-waker",
"bytes",
"fnv",
"futures-core",
"futures-sink",
"http",
"indexmap",
"parking_lot",
"slab",
"smallvec",
"tokio",
"tokio-util",
]
[[package]] [[package]]
name = "httparse" name = "httparse"
version = "1.10.1" version = "1.10.1"
@@ -2119,6 +2323,15 @@ version = "1.70.2"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695"
[[package]]
name = "itertools"
version = "0.13.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "413ee7dfc52ee1a4949ceeb7dbc8a33f2d6c088194d9f922fb8318faf1f01186"
dependencies = [
"either",
]
[[package]] [[package]]
name = "itertools" name = "itertools"
version = "0.14.0" version = "0.14.0"
@@ -2229,6 +2442,16 @@ version = "0.2.183"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d" checksum = "b5b646652bf6661599e1da8901b3b9522896f01e736bad5f723fe7a3a27f899d"
[[package]]
name = "libloading"
version = "0.8.9"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55"
dependencies = [
"cfg-if",
"windows-link",
]
[[package]] [[package]]
name = "libm" name = "libm"
version = "0.2.16" version = "0.2.16"
@@ -2565,6 +2788,17 @@ version = "1.70.2"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
[[package]]
name = "openssl-macros"
version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a948666b637a0f465e8564c73e89d4dde00d72d4d473cc972f390fc3dcee7d9c"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]] [[package]]
name = "openssl-probe" name = "openssl-probe"
version = "0.1.6" version = "0.1.6"
@@ -3000,7 +3234,7 @@ dependencies = [
"compact_str", "compact_str",
"hashbrown 0.16.1", "hashbrown 0.16.1",
"indoc", "indoc",
"itertools", "itertools 0.14.0",
"kasuari", "kasuari",
"lru", "lru",
"strum", "strum",
@@ -3052,7 +3286,7 @@ dependencies = [
"hashbrown 0.16.1", "hashbrown 0.16.1",
"indoc", "indoc",
"instability", "instability",
"itertools", "itertools 0.14.0",
"line-clipping", "line-clipping",
"ratatui-core", "ratatui-core",
"strum", "strum",
@@ -3392,6 +3626,17 @@ dependencies = [
"windows-sys 0.61.2", "windows-sys 0.61.2",
] ]
[[package]]
name = "schnellru"
version = "0.2.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "356285bbf17bea63d9e52e96bd18f039672ac92b55b8cb997d6162a2a37d1649"
dependencies = [
"ahash",
"cfg-if",
"hashbrown 0.13.2",
]
[[package]] [[package]]
name = "scopeguard" name = "scopeguard"
version = "1.2.0" version = "1.2.0"
@@ -4205,6 +4450,16 @@ dependencies = [
"windows-sys 0.61.2", "windows-sys 0.61.2",
] ]
[[package]]
name = "tokio-boring2"
version = "5.0.0-alpha.13"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0f81df1210d791f31d72d840de8fbd80b9c3cb324956523048b1413e2bd55756"
dependencies = [
"boring2",
"tokio",
]
[[package]] [[package]]
name = "tokio-macros" name = "tokio-macros"
version = "2.6.1" version = "2.6.1"
@@ -4236,6 +4491,18 @@ dependencies = [
"tokio", "tokio",
] ]
[[package]]
name = "tokio-socks"
version = "0.5.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0d4770b8024672c1101b3f6733eab95b18007dbe0847a8afe341fcf79e06043f"
dependencies = [
"either",
"futures-util",
"thiserror 1.0.69",
"tokio",
]
[[package]] [[package]]
name = "tokio-stream" name = "tokio-stream"
version = "0.1.18" version = "0.1.18"
@@ -4510,6 +4777,26 @@ dependencies = [
"utf-8", "utf-8",
] ]
[[package]]
name = "typed-builder"
version = "0.23.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "31aa81521b70f94402501d848ccc0ecaa8f93c8eb6999eb9747e72287757ffda"
dependencies = [
"typed-builder-macro",
]
[[package]]
name = "typed-builder-macro"
version = "0.23.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "076a02dc54dd46795c2e9c8282ed40bcfb1e22747e955de9389a1de28190fb26"
dependencies = [
"proc-macro2",
"quote",
"syn 2.0.117",
]
[[package]] [[package]]
name = "typenum" name = "typenum"
version = "1.19.0" version = "1.19.0"
@@ -4567,7 +4854,7 @@ version = "2.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "16b380a1238663e5f8a691f9039c73e1cdae598a30e9855f541d29b08b53e9a5" checksum = "16b380a1238663e5f8a691f9039c73e1cdae598a30e9855f541d29b08b53e9a5"
dependencies = [ dependencies = [
"itertools", "itertools 0.14.0",
"unicode-segmentation", "unicode-segmentation",
"unicode-width", "unicode-width",
] ]
@@ -4832,6 +5119,15 @@ dependencies = [
"wasm-bindgen", "wasm-bindgen",
] ]
[[package]]
name = "webpki-root-certs"
version = "1.0.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f31141ce3fc3e300ae89b78c0dd67f9708061d1d2eda54b8209346fd6be9a92c"
dependencies = [
"rustls-pki-types",
]
[[package]] [[package]]
name = "webpki-roots" name = "webpki-roots"
version = "0.26.11" version = "0.26.11"
@@ -5341,6 +5637,56 @@ dependencies = [
"wasmparser", "wasmparser",
] ]
[[package]]
name = "wreq"
version = "6.0.0-rc.28"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f79937f6c4df65b3f6f78715b9de2977afe9ee3b3436483c7949a24511e25935"
dependencies = [
"ahash",
"boring2",
"brotli",
"bytes",
"flate2",
"futures-channel",
"futures-util",
"http",
"http-body",
"http-body-util",
"http2",
"httparse",
"ipnet",
"libc",
"percent-encoding",
"pin-project-lite",
"schnellru",
"serde",
"serde_json",
"smallvec",
"socket2 0.6.3",
"sync_wrapper",
"tokio",
"tokio-boring2",
"tokio-socks",
"tokio-tungstenite 0.28.0",
"tokio-util",
"tower",
"url",
"want",
"webpki-root-certs",
"zstd",
]
[[package]]
name = "wreq-util"
version = "3.0.0-rc.10"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6c6bbe24d28beb9ceb58b514bd6a613c759d3b706f768b9d2950d5d35b543c04"
dependencies = [
"typed-builder",
"wreq",
]
[[package]] [[package]]
name = "writeable" name = "writeable"
version = "0.6.2" version = "0.6.2"
@@ -5482,3 +5828,31 @@ name = "zmij"
version = "1.0.21" version = "1.0.21"
source = "registry+https://github.com/rust-lang/crates.io-index" source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa"
[[package]]
name = "zstd"
version = "0.13.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e91ee311a569c327171651566e07972200e76fcfe2242a4fa446149a3881c08a"
dependencies = [
"zstd-safe",
]
[[package]]
name = "zstd-safe"
version = "7.2.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8f49c4d5f0abb602a93fb8736af2a4f4dd9512e36f7f570d66e65ff867ed3b9d"
dependencies = [
"zstd-sys",
]
[[package]]
name = "zstd-sys"
version = "2.0.16+zstd.1.5.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "91e19ebc2adc8f83e43039e79776e3fda8ca919132d68a1fed6a5faca2683748"
dependencies = [
"cc",
"pkg-config",
]

View File

@@ -6,7 +6,9 @@ members = [
"crates/aether-ai-serving", "crates/aether-ai-serving",
"crates/aether-pool-core", "crates/aether-pool-core",
"crates/aether-provider-pool", "crates/aether-provider-pool",
"crates/aether-routing-core",
"crates/aether-data-contracts", "crates/aether-data-contracts",
"crates/aether-data-query",
"crates/aether-data-schema", "crates/aether-data-schema",
"crates/aether-dispatch-core", "crates/aether-dispatch-core",
"crates/aether-cache", "crates/aether-cache",
@@ -41,7 +43,9 @@ aether-ai-formats = { path = "crates/aether-ai-formats" }
aether-ai-serving = { path = "crates/aether-ai-serving" } aether-ai-serving = { path = "crates/aether-ai-serving" }
aether-pool-core = { path = "crates/aether-pool-core" } aether-pool-core = { path = "crates/aether-pool-core" }
aether-provider-pool = { path = "crates/aether-provider-pool" } aether-provider-pool = { path = "crates/aether-provider-pool" }
aether-routing-core = { path = "crates/aether-routing-core" }
aether-data-contracts = { path = "crates/aether-data-contracts" } aether-data-contracts = { path = "crates/aether-data-contracts" }
aether-data-query = { path = "crates/aether-data-query" }
aether-data-schema = { path = "crates/aether-data-schema" } aether-data-schema = { path = "crates/aether-data-schema" }
aether-dispatch-core = { path = "crates/aether-dispatch-core" } aether-dispatch-core = { path = "crates/aether-dispatch-core" }
aether-cache = { path = "crates/aether-cache" } aether-cache = { path = "crates/aether-cache" }
@@ -94,6 +98,8 @@ tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
uuid = { version = "1", features = ["serde", "v4", "v5"] } uuid = { version = "1", features = ["serde", "v4", "v5"] }
webpki-roots = "0.26" webpki-roots = "0.26"
wreq = { version = "6.0.0-rc.28", default-features = false, features = ["json", "stream", "socks", "webpki-roots", "ws"] }
wreq-util = "3.0.0-rc.10"
url = "2" url = "2"
[profile.dev] [profile.dev]

View File

@@ -1,11 +1,14 @@
# syntax=docker/dockerfile:1 # syntax=docker/dockerfile:1
# Aether 运行镜像Rust gateway 直接服务 API + 前端静态文件(国内镜像源版本) # Aether 运行镜像Rust gateway 直接服务 API + 前端静态文件(国内镜像源版本)
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest . # 构建命令: docker build --build-arg AETHER_BUILD_VERSION=v0.7.2 -f Dockerfile.app.local -t aether-app:latest .
ARG RUST_VERSION=1.95.0 ARG RUST_VERSION=1.95.0
# ==================== 前端构建 ==================== # ==================== 前端构建 ====================
FROM node:22-slim AS frontend-builder FROM node:22-slim AS frontend-builder
ARG AETHER_BUILD_VERSION
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION}
WORKDIR /app/frontend WORKDIR /app/frontend
COPY frontend/package*.json ./ COPY frontend/package*.json ./
RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \ RUN --mount=type=cache,id=aether-npm-cache,target=/root/.npm,sharing=locked \
@@ -44,6 +47,9 @@ COPY crates/ ./crates/
RUN cargo chef prepare --recipe-path recipe.json RUN cargo chef prepare --recipe-path recipe.json
FROM gateway-base AS gateway-builder FROM gateway-base AS gateway-builder
ARG AETHER_BUILD_VERSION
ENV AETHER_BUILD_VERSION=${AETHER_BUILD_VERSION} \
AETHER_VERSION=${AETHER_BUILD_VERSION}
COPY --from=gateway-planner /build/recipe.json ./recipe.json COPY --from=gateway-planner /build/recipe.json ./recipe.json
RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \ RUN --mount=type=cache,id=aether-cargo-registry,target=/usr/local/cargo/registry,sharing=locked \
--mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \ --mount=type=cache,id=aether-cargo-git,target=/usr/local/cargo/git,sharing=locked \

View File

@@ -48,14 +48,14 @@ cp .env.example .env
./generate_keys.sh ./generate_keys.sh
# 编辑 .env 设置 ADMIN_PASSWORD # 编辑 .env 设置 ADMIN_PASSWORD
# 3. 首次部署 / 更新 (从以下数据库、内存策略任选其一) # 3. 首次部署 / 更新 (从以下部署形态任选其一)
# Postgres + Redis (适用于企业或多人使用) # Postgres + Redis (适用于企业或多人使用)
docker compose pull && docker compose up -d docker compose pull && docker compose up -d
# 仅SQLite (适用于个人用户或朋友分享) # Single Node (适用于个人用户或朋友分享)
docker compose -f docker-compose.sqlite.yml pull && docker compose -f docker-compose.sqlite.yml up -d docker compose -f docker-compose.single-node.yml pull && docker compose -f docker-compose.single-node.yml up -d
``` ```
### 一键安装(可选部署方式 Linux: systemd; Mac: launchd ### 一键安装(默认 Single NodeLinux systemd / macOS launchd + SQLite
```bash ```bash
cd Aether && cd Aether cd Aether && cd Aether

View File

@@ -23,6 +23,7 @@ aether-oauth.workspace = true
aether-pool-core.workspace = true aether-pool-core.workspace = true
aether-provider-pool.workspace = true aether-provider-pool.workspace = true
aether-provider-transport.workspace = true aether-provider-transport.workspace = true
aether-routing-core.workspace = true
aether-scheduler-core.workspace = true aether-scheduler-core.workspace = true
aether-runtime.workspace = true aether-runtime.workspace = true
aether-runtime-state.workspace = true aether-runtime-state.workspace = true
@@ -64,6 +65,8 @@ tracing.workspace = true
url.workspace = true url.workspace = true
uuid.workspace = true uuid.workspace = true
webpki-roots.workspace = true webpki-roots.workspace = true
wreq.workspace = true
wreq-util.workspace = true
[target.'cfg(not(target_env = "msvc"))'.dependencies] [target.'cfg(not(target_env = "msvc"))'.dependencies]
tikv-jemallocator = "0.6" tikv-jemallocator = "0.6"

View File

@@ -2,14 +2,20 @@ use std::env;
use std::process::Command; use std::process::Command;
fn main() { fn main() {
println!("cargo:rerun-if-env-changed=AETHER_BUILD_VERSION");
println!("cargo:rerun-if-env-changed=AETHER_VERSION"); println!("cargo:rerun-if-env-changed=AETHER_VERSION");
println!("cargo:rerun-if-env-changed=GITHUB_REF_NAME"); println!("cargo:rerun-if-env-changed=GITHUB_REF_NAME");
println!("cargo:rerun-if-changed=../../.git/HEAD"); println!("cargo:rerun-if-changed=../../.git/HEAD");
let package_version = env::var("CARGO_PKG_VERSION").unwrap_or_else(|_| "unknown".to_string()); let package_version = env::var("CARGO_PKG_VERSION").unwrap_or_else(|_| "unknown".to_string());
let version = env::var("AETHER_VERSION") let version = env::var("AETHER_BUILD_VERSION")
.ok() .ok()
.filter(|value| !value.trim().is_empty()) .filter(|value| !value.trim().is_empty())
.or_else(|| {
env::var("AETHER_VERSION")
.ok()
.filter(|value| !value.trim().is_empty())
})
.or_else(|| { .or_else(|| {
env::var("GITHUB_REF_NAME") env::var("GITHUB_REF_NAME")
.ok() .ok()

View File

@@ -41,26 +41,30 @@ pub(crate) use crate::ai_serving::{
AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt, AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt,
}; };
pub(crate) use aether_ai_formats::api::{ pub(crate) use aether_ai_formats::api::{
build_core_error_body_for_client_format, core_error_background_report_kind, build_core_error_body_for_client_format, convert_standard_chat_response,
core_error_default_client_api_format, core_success_background_report_kind, core_error_background_report_kind, core_error_default_client_api_format,
encode_kiro_sse_events, implicit_sync_finalize_report_kind, is_core_error_finalize_kind, core_success_background_report_kind, encode_kiro_sse_events,
implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
normalize_provider_private_report_context, normalize_provider_private_response_value, normalize_provider_private_report_context, normalize_provider_private_response_value,
provider_private_response_allows_sync_finalize, resolve_claude_stream_spec, provider_private_response_allows_sync_finalize, resolve_claude_stream_spec,
resolve_claude_sync_spec, resolve_gemini_stream_spec, resolve_gemini_sync_spec, resolve_claude_sync_spec, resolve_gemini_stream_spec, resolve_gemini_sync_spec,
resolve_local_image_stream_spec, resolve_local_image_sync_spec, resolve_local_image_stream_spec, resolve_local_image_sync_spec,
resolve_local_same_format_stream_spec, resolve_local_same_format_sync_spec, resolve_local_same_format_stream_spec, resolve_local_same_format_sync_spec,
resolve_openai_embedding_sync_spec, AiControlPlanRequest, ExecutionRuntimeAuthContext, resolve_openai_embedding_sync_spec, sanitize_request_path_and_query, AiControlPlanRequest,
LocalCoreSyncErrorKind, LocalOpenAiImageSpec, LocalSameFormatProviderFamily, CanonicalContentPart, CanonicalStreamEvent, CanonicalStreamFrame, ClaudeClientEmitter,
LocalSameFormatProviderSpec, LocalStandardSourceFamily, LocalStandardSourceMode, ExecutionRuntimeAuthContext, LocalCoreSyncErrorKind, LocalOpenAiImageSpec,
LocalStandardSpec, StreamingStandardTerminalObserver, EXECUTION_RUNTIME_STREAM_DECISION_ACTION, LocalSameFormatProviderFamily, LocalSameFormatProviderSpec, LocalStandardSourceFamily,
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_EMBEDDING_SYNC_PLAN_KIND, LocalStandardSourceMode, LocalStandardSpec, OpenAIChatClientEmitter,
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OpenAIResponsesClientEmitter, StreamingStandardTerminalObserver,
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND, EXECUTION_RUNTIME_STREAM_DECISION_ACTION, EXECUTION_RUNTIME_SYNC_DECISION_ACTION,
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND, GEMINI_EMBEDDING_SYNC_PLAN_KIND, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND, OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND, OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
}; };
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
pub(crate) fn parse_direct_request_body( pub(crate) fn parse_direct_request_body(
parts: &http::request::Parts, parts: &http::request::Parts,

View File

@@ -180,7 +180,9 @@ pub(crate) fn resolve_local_decision_execution_runtime_auth_context(
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
) -> Option<ExecutionRuntimeAuthContext> { ) -> Option<ExecutionRuntimeAuthContext> {
resolve_decision_execution_runtime_auth_context(decision).filter(|auth_context| { resolve_decision_execution_runtime_auth_context(decision).filter(|auth_context| {
!auth_context.user_id.trim().is_empty() && !auth_context.api_key_id.trim().is_empty() auth_context.access_allowed
&& !auth_context.user_id.trim().is_empty()
&& !auth_context.api_key_id.trim().is_empty()
}) })
} }

View File

@@ -7,7 +7,13 @@ use aether_ai_serving::{
AiCandidatePreselectionOutcome, AiSkippedCandidatePersistencePort, AiCandidatePreselectionOutcome, AiSkippedCandidatePersistencePort,
}; };
use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem}; use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem};
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate}; use aether_routing_core::{
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
RoutingCandidateTrace, RoutingDecisionTrace,
};
use aether_scheduler_core::{
ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate, SchedulerRankingOutcome,
};
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::Value; use serde_json::Value;
use std::collections::VecDeque; use std::collections::VecDeque;
@@ -19,6 +25,7 @@ use tracing::warn;
use uuid::Uuid; use uuid::Uuid;
use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_at_epoch; use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_at_epoch;
use crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy;
use crate::ai_serving::planner::candidate_resolution::{ use crate::ai_serving::planner::candidate_resolution::{
resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate, resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate,
LocalExecutionCandidateKind, SkippedLocalExecutionCandidate, LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
@@ -36,7 +43,7 @@ use crate::dispatch::refs::dispatch_ref_for_local_candidate;
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value; use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity}; use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::scheduler::candidate::API_KEY_CONCURRENCY_LIMIT_SKIP_REASON; use crate::scheduler::candidate::API_KEY_CONCURRENCY_LIMIT_SKIP_REASON;
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode}; use crate::scheduler::config::SchedulerSchedulingMode;
use crate::{AppState, GatewayError}; use crate::{AppState, GatewayError};
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100; const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
@@ -189,6 +196,7 @@ struct GatewayLocalCandidateMaterializationPort<'a, F, G> {
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&'a ClientSessionAffinity>, client_session_affinity: Option<&'a ClientSessionAffinity>,
required_capabilities: Option<&'a Value>, required_capabilities: Option<&'a Value>,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
sticky_session_token: Option<&'a str>, sticky_session_token: Option<&'a str>,
request_auth_channel: Option<&'a str>, request_auth_channel: Option<&'a str>,
persistence_policy: LocalCandidatePersistencePolicy<'a>, persistence_policy: LocalCandidatePersistencePolicy<'a>,
@@ -243,6 +251,7 @@ where
self.auth_snapshot, self.auth_snapshot,
self.client_session_affinity, self.client_session_affinity,
self.required_capabilities, self.required_capabilities,
self.routing_policy,
self.sticky_session_token, self.sticky_session_token,
self.request_auth_channel, self.request_auth_channel,
self.resolution_mode, self.resolution_mode,
@@ -281,6 +290,8 @@ where
.skipped .skipped
.record_runtime_miss_diagnostic, .record_runtime_miss_diagnostic,
candidates, candidates,
self.routing_policy,
self.client_api_format,
self.sticky_session_token, self.sticky_session_token,
self.requested_model, self.requested_model,
self.request_auth_channel, self.request_auth_channel,
@@ -294,6 +305,12 @@ where
starting_candidate_index: u32, starting_candidate_index: u32,
skipped_candidates: Vec<Self::Skipped>, skipped_candidates: Vec<Self::Skipped>,
) -> Result<(), Self::Error> { ) -> Result<(), Self::Error> {
let skipped_candidates = attach_routing_trace_to_skipped_candidates(
self.routing_policy,
self.client_api_format,
starting_candidate_index,
skipped_candidates,
);
persist_skipped_local_execution_candidates_with_context( persist_skipped_local_execution_candidates_with_context(
self.state.app(), self.state.app(),
self.trace_id, self.trace_id,
@@ -432,6 +449,7 @@ pub(crate) async fn materialize_local_execution_candidates_with_serving<F, G>(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&Value>, required_capabilities: Option<&Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
persistence_policy: LocalCandidatePersistencePolicy<'_>, persistence_policy: LocalCandidatePersistencePolicy<'_>,
@@ -445,7 +463,8 @@ where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync, F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync, G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
{ {
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await; let scheduler_cache_affinity_enabled =
scheduler_cache_affinity_enabled(state, routing_policy).await;
let port = GatewayLocalCandidateMaterializationPort { let port = GatewayLocalCandidateMaterializationPort {
state, state,
trace_id, trace_id,
@@ -454,6 +473,7 @@ where
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
sticky_session_token, sticky_session_token,
request_auth_channel, request_auth_channel,
persistence_policy, persistence_policy,
@@ -478,6 +498,7 @@ pub(crate) async fn build_local_execution_candidate_attempt_source_with_serving<
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&Value>, required_capabilities: Option<&Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
persistence_policy: LocalCandidatePersistencePolicy<'_>, persistence_policy: LocalCandidatePersistencePolicy<'_>,
@@ -491,7 +512,8 @@ where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync, F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync, G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
{ {
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await; let scheduler_cache_affinity_enabled =
scheduler_cache_affinity_enabled(state, routing_policy).await;
let _ = build_available_extra_data; let _ = build_available_extra_data;
let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates( let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates(
state, state,
@@ -501,6 +523,7 @@ where
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
sticky_session_token, sticky_session_token,
request_auth_channel, request_auth_channel,
resolution_mode, resolution_mode,
@@ -529,7 +552,12 @@ where
trace_id, trace_id,
persistence_policy.skipped, persistence_policy.skipped,
u32::try_from(candidates.len()).unwrap_or(u32::MAX), u32::try_from(candidates.len()).unwrap_or(u32::MAX),
skipped_candidates, attach_routing_trace_to_skipped_candidates(
routing_policy,
client_api_format,
u32::try_from(candidates.len()).unwrap_or(u32::MAX),
skipped_candidates,
),
) )
.await; .await;
@@ -542,6 +570,7 @@ where
sticky_session_token, sticky_session_token,
requested_model, requested_model,
request_auth_channel, request_auth_channel,
routing_policy,
); );
( (
@@ -559,6 +588,7 @@ fn build_logical_candidate_items<'a>(
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
requested_model: Option<&str>, requested_model: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> (VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>, u32) { ) -> (VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>, u32) {
let mut items = VecDeque::new(); let mut items = VecDeque::new();
let mut next_candidate_index = starting_candidate_index; let mut next_candidate_index = starting_candidate_index;
@@ -578,12 +608,13 @@ fn build_logical_candidate_items<'a>(
} }
} }
LocalExecutionCandidateKind::PoolGroup => { LocalExecutionCandidateKind::PoolGroup => {
let cursor = PoolKeyCursor::new( let cursor = PoolKeyCursor::new_with_routing_policy(
state, state,
candidate, candidate,
sticky_session_token, sticky_session_token,
requested_model, requested_model,
request_auth_channel, request_auth_channel,
routing_policy,
); );
let cursor = if let Some(trace_id) = trace_id { let cursor = if let Some(trace_id) = trace_id {
cursor.with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic) cursor.with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic)
@@ -615,6 +646,7 @@ pub(crate) async fn build_lazy_requested_model_execution_candidate_attempt_sourc
auth_snapshot: &GatewayAuthApiKeySnapshot, auth_snapshot: &GatewayAuthApiKeySnapshot,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&Value>, required_capabilities: Option<&Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
persistence_policy: LocalCandidatePersistencePolicy<'_>, persistence_policy: LocalCandidatePersistencePolicy<'_>,
@@ -628,7 +660,8 @@ where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync + 'a, F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync + 'a,
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync + 'a, G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync + 'a,
{ {
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await; let scheduler_cache_affinity_enabled =
scheduler_cache_affinity_enabled(state, routing_policy).await;
let _ = build_available_extra_data; let _ = build_available_extra_data;
let decorate_skipped_candidate = Arc::new(decorate_skipped_candidate); let decorate_skipped_candidate = Arc::new(decorate_skipped_candidate);
let record_runtime_miss_diagnostic = persistence_policy.skipped.record_runtime_miss_diagnostic; let record_runtime_miss_diagnostic = persistence_policy.skipped.record_runtime_miss_diagnostic;
@@ -639,6 +672,7 @@ where
require_streaming, require_streaming,
required_capabilities, required_capabilities,
auth_snapshot, auth_snapshot,
routing_policy,
client_session_affinity, client_session_affinity,
use_api_format_alias_match, use_api_format_alias_match,
key_mode, key_mode,
@@ -652,6 +686,7 @@ where
auth_snapshot: auth_snapshot.clone(), auth_snapshot: auth_snapshot.clone(),
client_session_affinity: client_session_affinity.cloned(), client_session_affinity: client_session_affinity.cloned(),
required_capabilities: required_capabilities.cloned(), required_capabilities: required_capabilities.cloned(),
routing_policy: routing_policy.cloned(),
sticky_session_token: sticky_session_token.map(str::to_string), sticky_session_token: sticky_session_token.map(str::to_string),
request_auth_channel: request_auth_channel.map(str::to_string), request_auth_channel: request_auth_channel.map(str::to_string),
skipped_user_id: persistence_policy.skipped.user_id.to_string(), skipped_user_id: persistence_policy.skipped.user_id.to_string(),
@@ -693,6 +728,7 @@ struct RequestedModelAttemptPageCursor<'a> {
auth_snapshot: GatewayAuthApiKeySnapshot, auth_snapshot: GatewayAuthApiKeySnapshot,
client_session_affinity: Option<ClientSessionAffinity>, client_session_affinity: Option<ClientSessionAffinity>,
required_capabilities: Option<Value>, required_capabilities: Option<Value>,
routing_policy: Option<ResolvedRoutingPolicy>,
sticky_session_token: Option<String>, sticky_session_token: Option<String>,
request_auth_channel: Option<String>, request_auth_channel: Option<String>,
skipped_user_id: String, skipped_user_id: String,
@@ -756,6 +792,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
Some(&self.auth_snapshot), Some(&self.auth_snapshot),
self.client_session_affinity.as_ref(), self.client_session_affinity.as_ref(),
self.required_capabilities.as_ref(), self.required_capabilities.as_ref(),
self.routing_policy.as_ref(),
self.sticky_session_token.as_deref(), self.sticky_session_token.as_deref(),
self.request_auth_channel.as_deref(), self.request_auth_channel.as_deref(),
self.resolution_mode, self.resolution_mode,
@@ -794,6 +831,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
self.sticky_session_token.as_deref(), self.sticky_session_token.as_deref(),
Some(&self.requested_model), Some(&self.requested_model),
self.request_auth_channel.as_deref(), self.request_auth_channel.as_deref(),
self.routing_policy.as_ref(),
); );
self.next_candidate_index = next_candidate_index self.next_candidate_index = next_candidate_index
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX)); .saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
@@ -814,7 +852,12 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
&self.trace_id, &self.trace_id,
skipped_persistence, skipped_persistence,
skipped_starting_candidate_index, skipped_starting_candidate_index,
skipped_candidates, attach_routing_trace_to_skipped_candidates(
self.routing_policy.as_ref(),
&self.client_api_format,
skipped_starting_candidate_index,
skipped_candidates,
),
) )
.await; .await;
} }
@@ -858,7 +901,12 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
&self.trace_id, &self.trace_id,
skipped_persistence, skipped_persistence,
self.next_candidate_index, self.next_candidate_index,
skipped_candidates, attach_routing_trace_to_skipped_candidates(
self.routing_policy.as_ref(),
&self.client_api_format,
self.next_candidate_index,
skipped_candidates,
),
) )
.await; .await;
self.next_candidate_index = self self.next_candidate_index = self
@@ -925,19 +973,14 @@ async fn pop_attempt_from_items(
} }
} }
async fn scheduler_cache_affinity_enabled(state: PlannerAppState<'_>) -> bool { async fn scheduler_cache_affinity_enabled(
match read_scheduler_ordering_config(state.app()).await { state: PlannerAppState<'_>,
Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity, routing_policy: Option<&ResolvedRoutingPolicy>,
Err(error) => { ) -> bool {
warn!( scheduler_ordering_config_for_routing_policy(state, routing_policy)
event_name = "planner_scheduler_affinity_config_load_failed", .await
log_type = "event", .scheduling_mode
error = ?error, == SchedulerSchedulingMode::CacheAffinity
"failed to load scheduler config while checking cache affinity mode"
);
SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity
}
}
} }
pub(crate) fn remember_first_local_candidate_affinity( pub(crate) fn remember_first_local_candidate_affinity(
@@ -1038,6 +1081,8 @@ async fn materialize_logical_local_execution_candidate_attempts<F>(
context: LocalAvailableCandidatePersistenceContext<'_>, context: LocalAvailableCandidatePersistenceContext<'_>,
record_runtime_miss_diagnostic: bool, record_runtime_miss_diagnostic: bool,
candidates: Vec<EligibleLocalExecutionCandidate>, candidates: Vec<EligibleLocalExecutionCandidate>,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_api_format: &str,
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
requested_model: Option<&str>, requested_model: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
@@ -1059,18 +1104,21 @@ where
context, context,
candidate, candidate,
candidate_index, candidate_index,
routing_policy,
client_api_format,
build_extra_data, build_extra_data,
) )
.await, .await,
); );
} }
LocalExecutionCandidateKind::PoolGroup => { LocalExecutionCandidateKind::PoolGroup => {
let mut cursor = PoolKeyCursor::new( let mut cursor = PoolKeyCursor::new_with_routing_policy(
state, state,
candidate, candidate,
sticky_session_token, sticky_session_token,
requested_model, requested_model,
request_auth_channel, request_auth_channel,
routing_policy,
) )
.with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic); .with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic);
let attempt_count_before_pool = attempts.len(); let attempt_count_before_pool = attempts.len();
@@ -1097,6 +1145,8 @@ async fn persist_available_local_execution_candidate_at_index<F>(
context: LocalAvailableCandidatePersistenceContext<'_>, context: LocalAvailableCandidatePersistenceContext<'_>,
candidate: EligibleLocalExecutionCandidate, candidate: EligibleLocalExecutionCandidate,
candidate_index: u32, candidate_index: u32,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_api_format: &str,
build_extra_data: &F, build_extra_data: &F,
) -> Vec<LocalExecutionCandidateAttempt> ) -> Vec<LocalExecutionCandidateAttempt>
where where
@@ -1107,6 +1157,16 @@ where
available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data), available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data),
candidate.ranking.as_ref(), candidate.ranking.as_ref(),
); );
let extra_data = attach_routing_trace_to_extra_data(
routing_policy,
client_api_format,
&candidate.candidate,
candidate.kind,
candidate.ranking.as_ref(),
None,
Some(candidate_index),
extra_data,
);
let should_persist = should_persist_available_local_candidate(&candidate); let should_persist = should_persist_available_local_candidate(&candidate);
let mut attempts = Vec::with_capacity(attempt_slots as usize); let mut attempts = Vec::with_capacity(attempt_slots as usize);
let mut owned_candidate = Some(candidate); let mut owned_candidate = Some(candidate);
@@ -1190,6 +1250,160 @@ where
Some(Value::Object(object)) Some(Value::Object(object))
} }
fn attach_routing_trace_to_skipped_candidates(
routing_policy: Option<&ResolvedRoutingPolicy>,
client_api_format: &str,
starting_candidate_index: u32,
skipped_candidates: Vec<SkippedLocalExecutionCandidate>,
) -> Vec<SkippedLocalExecutionCandidate> {
skipped_candidates
.into_iter()
.enumerate()
.map(|(offset, skipped)| {
let selected_order =
starting_candidate_index.saturating_add(u32::try_from(offset).unwrap_or(u32::MAX));
attach_routing_trace_to_skipped_candidate(
routing_policy,
client_api_format,
selected_order,
skipped,
)
})
.collect()
}
fn attach_routing_trace_to_skipped_candidate(
routing_policy: Option<&ResolvedRoutingPolicy>,
client_api_format: &str,
selected_order: u32,
mut skipped_candidate: SkippedLocalExecutionCandidate,
) -> SkippedLocalExecutionCandidate {
let kind = if skipped_candidate
.transport
.as_ref()
.is_some_and(|transport| {
admin_provider_pool_config_from_config_value(transport.provider.config.as_ref())
.is_some()
}) {
LocalExecutionCandidateKind::PoolGroup
} else {
LocalExecutionCandidateKind::SingleKey
};
skipped_candidate.extra_data = attach_routing_trace_to_extra_data(
routing_policy,
client_api_format,
&skipped_candidate.candidate,
kind,
skipped_candidate.ranking.as_ref(),
Some(skipped_candidate.skip_reason),
Some(selected_order),
skipped_candidate.extra_data,
);
skipped_candidate
}
#[allow(clippy::too_many_arguments)]
fn attach_routing_trace_to_extra_data(
routing_policy: Option<&ResolvedRoutingPolicy>,
client_api_format: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
kind: LocalExecutionCandidateKind,
ranking: Option<&SchedulerRankingOutcome>,
skip_reason: Option<&'static str>,
selected_order: Option<u32>,
extra_data: Option<Value>,
) -> Option<Value> {
let Some(policy) = routing_policy else {
return extra_data;
};
let routing_trace = routing_trace_for_candidate(
policy,
client_api_format,
candidate,
kind,
ranking,
skip_reason,
selected_order,
);
Some(merge_routing_trace_into_extra_data(
extra_data,
routing_trace,
))
}
fn merge_routing_trace_into_extra_data(
extra_data: Option<Value>,
routing_trace: RoutingDecisionTrace,
) -> Value {
let mut object = match extra_data {
Some(Value::Object(object)) => object,
Some(value) => {
let mut object = serde_json::Map::new();
object.insert("extra".to_string(), value);
object
}
None => serde_json::Map::new(),
};
object.insert(
"routing_trace".to_string(),
serde_json::json!(routing_trace),
);
Value::Object(object)
}
fn routing_trace_for_candidate(
policy: &ResolvedRoutingPolicy,
client_api_format: &str,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
kind: LocalExecutionCandidateKind,
ranking: Option<&SchedulerRankingOutcome>,
skip_reason: Option<&'static str>,
selected_order: Option<u32>,
) -> RoutingDecisionTrace {
let candidate_kind = routing_candidate_kind(kind);
let mut trace = crate::routing::build_routing_trace_seed(policy, client_api_format);
trace.global_candidates.push(RoutingCandidateTrace {
candidate_kind,
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
model_id: candidate.model_id.clone(),
key_id: match candidate_kind {
CandidateKind::Provider => Some(candidate.key_id.clone()),
CandidateKind::PoolGroup => None,
},
ranking_vector: rank_vector_for_candidate(
&policy.ranking_overlay,
&RoutingCandidateFacts {
candidate_kind,
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
model_id: candidate.model_id.clone(),
key_id: match candidate_kind {
CandidateKind::Provider => Some(candidate.key_id.clone()),
CandidateKind::PoolGroup => None,
},
provider_priority: candidate.provider_priority,
key_priority: candidate
.key_global_priority_for_format
.unwrap_or(candidate.key_internal_priority),
},
),
skip_reason: skip_reason.map(str::to_string),
selected_order,
});
if let Some(ranking) = ranking {
trace.runtime_facts.cache_affinity_hit = ranking.promoted_by == Some("cached_affinity");
}
trace
}
fn routing_candidate_kind(kind: LocalExecutionCandidateKind) -> CandidateKind {
match kind {
LocalExecutionCandidateKind::SingleKey => CandidateKind::Provider,
LocalExecutionCandidateKind::PoolGroup => CandidateKind::PoolGroup,
}
}
fn dispatch_sequence_from_attempts( fn dispatch_sequence_from_attempts(
attempts: Vec<LocalExecutionCandidateAttempt>, attempts: Vec<LocalExecutionCandidateAttempt>,
) -> DispatchSequence<LocalExecutionCandidateAttempt> { ) -> DispatchSequence<LocalExecutionCandidateAttempt> {
@@ -1644,6 +1858,7 @@ mod tests {
auth_snapshot: Some(&auth_snapshot), auth_snapshot: Some(&auth_snapshot),
client_session_affinity: None, client_session_affinity: None,
required_capabilities: None, required_capabilities: None,
routing_policy: None,
sticky_session_token: None, sticky_session_token: None,
request_auth_channel: None, request_auth_channel: None,
persistence_policy: LocalCandidatePersistencePolicy { persistence_policy: LocalCandidatePersistencePolicy {
@@ -1718,6 +1933,8 @@ mod tests {
false, false,
vec![pool_group, sample_eligible("normal-key", None)], vec![pool_group, sample_eligible("normal-key", None)],
None, None,
"openai:chat",
None,
Some("gpt-5"), Some("gpt-5"),
None, None,
&|_| None, &|_| None,

View File

@@ -5,6 +5,7 @@ use aether_ai_serving::{
AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig, AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig,
AiRankingSchedulingMode, AiRankingSchedulingMode,
}; };
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode};
use async_trait::async_trait; use async_trait::async_trait;
use tracing::warn; use tracing::warn;
@@ -16,12 +17,12 @@ use crate::scheduler::config::{
}; };
use aether_scheduler_core::{ use aether_scheduler_core::{
matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget, matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget,
SchedulerMinimalCandidateSelectionCandidate, SchedulerRankableCandidate, SchedulerMinimalCandidateSelectionCandidate, SchedulerPriorityMode, SchedulerRankableCandidate,
SchedulerRankingContext, SchedulerRankingOutcome, SchedulerRankingContext, SchedulerRankingOutcome,
}; };
use super::candidate_affinity_cache::read_cached_scheduler_affinity_target; use super::candidate_affinity_cache::read_cached_scheduler_affinity_target;
use super::candidate_resolution::EligibleLocalExecutionCandidate; use super::candidate_resolution::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
use super::candidate_transport_ranking_facts::{ use super::candidate_transport_ranking_facts::{
resolve_cached_transport_ranking_facts, CandidateTransportRankingFacts, resolve_cached_transport_ranking_facts, CandidateTransportRankingFacts,
}; };
@@ -33,6 +34,7 @@ struct GatewayLocalCandidateRankingPort<'a> {
client_session_affinity: Option<&'a ClientSessionAffinity>, client_session_affinity: Option<&'a ClientSessionAffinity>,
required_capabilities: Option<&'a serde_json::Value>, required_capabilities: Option<&'a serde_json::Value>,
ordering_config: SchedulerOrderingConfig, ordering_config: SchedulerOrderingConfig,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
} }
#[async_trait] #[async_trait]
@@ -89,8 +91,10 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> {
self.ordering_config, self.ordering_config,
) )
.await; .await;
let routing_overlaid_candidate =
routing_overlaid_candidate(self.routing_policy, candidate.kind, &candidate.candidate);
Ok(build_ai_rankable_candidate(AiRankableCandidateParts { Ok(build_ai_rankable_candidate(AiRankableCandidateParts {
candidate: &candidate.candidate, candidate: &routing_overlaid_candidate,
original_index, original_index,
normalized_client_api_format, normalized_client_api_format,
provider_api_format: candidate.provider_api_format.as_str(), provider_api_format: candidate.provider_api_format.as_str(),
@@ -122,8 +126,9 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> Vec<EligibleLocalExecutionCandidate> { ) -> Vec<EligibleLocalExecutionCandidate> {
let ordering_config = read_scheduler_ordering_config_or_default(state).await; let ordering_config = scheduler_ordering_config_for_routing_policy(state, routing_policy).await;
let port = GatewayLocalCandidateRankingPort { let port = GatewayLocalCandidateRankingPort {
state, state,
requested_model, requested_model,
@@ -131,6 +136,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
ordering_config, ordering_config,
routing_policy,
}; };
match run_ai_candidate_ranking(&port, candidates, normalized_client_api_format).await { match run_ai_candidate_ranking(&port, candidates, normalized_client_api_format).await {
@@ -189,6 +195,58 @@ fn ai_ranking_scheduling_mode(mode: SchedulerSchedulingMode) -> AiRankingSchedul
} }
} }
pub(crate) async fn scheduler_ordering_config_for_routing_policy(
state: PlannerAppState<'_>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> SchedulerOrderingConfig {
match routing_policy {
Some(policy) => scheduler_ordering_config_from_routing_policy(policy),
None => read_scheduler_ordering_config_or_default(state).await,
}
}
fn scheduler_ordering_config_from_routing_policy(
policy: &ResolvedRoutingPolicy,
) -> SchedulerOrderingConfig {
SchedulerOrderingConfig {
priority_mode: match policy.priority_mode {
RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider,
RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey,
},
scheduling_mode: match policy.scheduling_mode {
RoutingSchedulingMode::FixedOrder => SchedulerSchedulingMode::FixedOrder,
RoutingSchedulingMode::CacheAffinity => SchedulerSchedulingMode::CacheAffinity,
RoutingSchedulingMode::LoadBalance => SchedulerSchedulingMode::LoadBalance,
},
keep_priority_on_conversion: policy.keep_priority_on_conversion,
}
}
fn routing_overlaid_candidate(
routing_policy: Option<&ResolvedRoutingPolicy>,
kind: LocalExecutionCandidateKind,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> SchedulerMinimalCandidateSelectionCandidate {
let Some(policy) = routing_policy else {
return candidate.clone();
};
let mut overlaid = candidate.clone();
overlaid.provider_priority = policy
.ranking_overlay
.provider_priority_or_unspecified(candidate.provider_id.as_str());
let overlaid_key_priority = match kind {
LocalExecutionCandidateKind::SingleKey => policy
.ranking_overlay
.key_priority_or_unspecified(candidate.key_id.as_str()),
LocalExecutionCandidateKind::PoolGroup => policy
.ranking_overlay
.pool_priority_or_unspecified(candidate.provider_id.as_str()),
};
overlaid.key_internal_priority = overlaid_key_priority;
overlaid.key_global_priority_for_format = Some(overlaid_key_priority);
overlaid
}
async fn read_scheduler_ordering_config_or_default( async fn read_scheduler_ordering_config_or_default(
state: PlannerAppState<'_>, state: PlannerAppState<'_>,
) -> SchedulerOrderingConfig { ) -> SchedulerOrderingConfig {
@@ -304,6 +362,82 @@ mod tests {
} }
} }
#[test]
fn routing_policy_priorities_do_not_fall_back_to_candidate_priorities() {
let mut candidate = sample_candidate("endpoint-1", "key-1");
candidate.provider_priority = 7;
candidate.key_internal_priority = 3;
candidate.key_global_priority_for_format = Some(2);
let policy = aether_routing_core::ResolvedRoutingPolicy {
group_id: Some("group-1".to_string()),
group_version: Some(1),
selection_source: "system_default".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
keep_priority_on_conversion: false,
ranking_overlay: aether_routing_core::RankingOverlay::default(),
mutation_plan: Default::default(),
pool_policy_overrides: BTreeMap::new(),
matched_rules: Vec::new(),
};
let overlaid = super::routing_overlaid_candidate(
Some(&policy),
LocalExecutionCandidateKind::SingleKey,
&candidate,
);
assert_eq!(
overlaid.provider_priority,
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
);
assert_eq!(
overlaid.key_internal_priority,
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
);
assert_eq!(
overlaid.key_global_priority_for_format,
Some(aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED)
);
}
#[test]
fn routing_policy_uses_pool_priority_for_pool_group_global_key_slot() {
let mut candidate = sample_candidate("endpoint-1", "representative-key");
candidate.provider_priority = 7;
candidate.key_internal_priority = 3;
candidate.key_global_priority_for_format = Some(2);
let policy = aether_routing_core::ResolvedRoutingPolicy {
group_id: Some("group-1".to_string()),
group_version: Some(1),
selection_source: "system_default".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: aether_routing_core::RoutingSetPriorityMode::GlobalKey,
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
keep_priority_on_conversion: false,
ranking_overlay: aether_routing_core::RankingOverlay {
pool_priority_overrides: BTreeMap::from([("provider-1".to_string(), 4)]),
key_priority_overrides: BTreeMap::from([("representative-key".to_string(), 1)]),
..Default::default()
},
mutation_plan: Default::default(),
pool_policy_overrides: BTreeMap::new(),
matched_rules: Vec::new(),
};
let overlaid = super::routing_overlaid_candidate(
Some(&policy),
LocalExecutionCandidateKind::PoolGroup,
&candidate,
);
assert_eq!(overlaid.key_internal_priority, 4);
assert_eq!(overlaid.key_global_priority_for_format, Some(4));
}
fn sample_provider() -> StoredProviderCatalogProvider { fn sample_provider() -> StoredProviderCatalogProvider {
sample_provider_with_options("provider-1", false, 0) sample_provider_with_options("provider-1", false, 0)
} }
@@ -1062,6 +1196,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1141,6 +1276,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1216,6 +1352,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1282,6 +1419,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1364,6 +1502,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1438,6 +1577,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1515,6 +1655,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1610,6 +1751,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1713,6 +1855,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
) )
.await; .await;
@@ -1809,6 +1952,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
aether_ai_serving::AiCandidateResolutionMode::Standard, aether_ai_serving::AiCandidateResolutionMode::Standard,
) )
.await; .await;
@@ -1901,6 +2045,7 @@ mod tests {
None, None,
None, None,
None, None,
None,
aether_ai_serving::AiCandidateResolutionMode::Standard, aether_ai_serving::AiCandidateResolutionMode::Standard,
) )
.await; .await;

View File

@@ -4,6 +4,7 @@ use aether_ai_serving::{
run_ai_candidate_resolution, AiCandidateResolutionMode, AiCandidateResolutionPort, run_ai_candidate_resolution, AiCandidateResolutionMode, AiCandidateResolutionPort,
AiCandidateResolutionRequest, AiCandidateResolutionRequest,
}; };
use aether_routing_core::ResolvedRoutingPolicy;
use async_trait::async_trait; use async_trait::async_trait;
use std::convert::Infallible; use std::convert::Infallible;
use tracing::warn; use tracing::warn;
@@ -60,6 +61,7 @@ struct GatewayLocalCandidateResolutionPort<'a> {
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&'a ClientSessionAffinity>, client_session_affinity: Option<&'a ClientSessionAffinity>,
required_capabilities: Option<&'a serde_json::Value>, required_capabilities: Option<&'a serde_json::Value>,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
request_auth_channel: Option<&'a str>, request_auth_channel: Option<&'a str>,
} }
@@ -97,6 +99,11 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
transport: &Self::Transport, transport: &Self::Transport,
requested_model: Option<&str>, requested_model: Option<&str>,
) -> Option<&'static str> { ) -> Option<&'static str> {
if let Some(skip_reason) =
routing_policy_candidate_skip_reason(self.routing_policy, candidate, transport)
{
return Some(skip_reason);
}
if provider_transport_uses_pool(transport) { if provider_transport_uses_pool(transport) {
return pool_group_common_transport_skip_reason(candidate, transport); return pool_group_common_transport_skip_reason(candidate, transport);
} }
@@ -172,6 +179,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
self.auth_snapshot, self.auth_snapshot,
self.client_session_affinity, self.client_session_affinity,
self.required_capabilities, self.required_capabilities,
self.routing_policy,
) )
.await) .await)
} }
@@ -192,6 +200,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
_sticky_session_token: Option<&str>, _sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
) -> ( ) -> (
@@ -207,6 +216,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
None, None,
request_auth_channel, request_auth_channel,
AiCandidateResolutionMode::Standard, AiCandidateResolutionMode::Standard,
@@ -222,6 +232,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
_sticky_session_token: Option<&str>, _sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
) -> ( ) -> (
@@ -237,6 +248,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
None, None,
request_auth_channel, request_auth_channel,
AiCandidateResolutionMode::WithoutTransportPairGate, AiCandidateResolutionMode::WithoutTransportPairGate,
@@ -252,6 +264,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
_sticky_session_token: Option<&str>, _sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode, mode: AiCandidateResolutionMode,
@@ -267,6 +280,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
None, None,
request_auth_channel, request_auth_channel,
mode, mode,
@@ -283,6 +297,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
_sticky_session_token: Option<&str>, _sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode, mode: AiCandidateResolutionMode,
@@ -298,6 +313,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
None, None,
request_auth_channel, request_auth_channel,
mode, mode,
@@ -315,6 +331,7 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>, auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
routing_policy: Option<&ResolvedRoutingPolicy>,
_sticky_session_token: Option<&str>, _sticky_session_token: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
mode: AiCandidateResolutionMode, mode: AiCandidateResolutionMode,
@@ -330,6 +347,7 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
auth_snapshot, auth_snapshot,
client_session_affinity, client_session_affinity,
required_capabilities, required_capabilities,
routing_policy,
request_auth_channel, request_auth_channel,
}; };
@@ -369,6 +387,28 @@ fn provider_transport_uses_pool(transport: &GatewayProviderTransportSnapshot) ->
.is_some() .is_some()
} }
fn routing_policy_candidate_skip_reason(
routing_policy: Option<&ResolvedRoutingPolicy>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
) -> Option<&'static str> {
let policy = routing_policy?;
if !policy
.ranking_overlay
.provider_allowed(candidate.provider_id.as_str())
{
return Some("routing_profile_disallowed_provider");
}
if !provider_transport_uses_pool(transport)
&& !policy
.ranking_overlay
.key_allowed(candidate.key_id.as_str())
{
return Some("routing_profile_disallowed_key");
}
None
}
fn pool_group_common_transport_skip_reason( fn pool_group_common_transport_skip_reason(
candidate: &SchedulerMinimalCandidateSelectionCandidate, candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,

View File

@@ -2,6 +2,7 @@ use aether_ai_serving::{
run_ai_candidate_preselection, AiCandidatePreselectionOutcome, AiCandidatePreselectionPort, run_ai_candidate_preselection, AiCandidatePreselectionOutcome, AiCandidatePreselectionPort,
}; };
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow; use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
use aether_routing_core::ResolvedRoutingPolicy;
use aether_scheduler_core::{ use aether_scheduler_core::{
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format, enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
resolve_requested_global_model_name_with_model_directives, resolve_requested_global_model_name_with_model_directives,
@@ -35,6 +36,7 @@ struct GatewayLocalCandidatePreselectionPort<'a> {
require_streaming: bool, require_streaming: bool,
required_capabilities: Option<&'a serde_json::Value>, required_capabilities: Option<&'a serde_json::Value>,
auth_snapshot: &'a GatewayAuthApiKeySnapshot, auth_snapshot: &'a GatewayAuthApiKeySnapshot,
routing_policy: Option<&'a ResolvedRoutingPolicy>,
client_session_affinity: Option<&'a ClientSessionAffinity>, client_session_affinity: Option<&'a ClientSessionAffinity>,
use_api_format_alias_match: bool, use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode, key_mode: LocalCandidatePreselectionKeyMode,
@@ -100,13 +102,14 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
let enable_model_directives = self.model_directive_enabled_api_formats.contains( let enable_model_directives = self.model_directive_enabled_api_formats.contains(
&crate::ai_serving::normalize_api_format_alias(candidate_api_format), &crate::ai_serving::normalize_api_format_alias(candidate_api_format),
); );
matches_client_format routing_policy_allows_provider(self.routing_policy, candidate)
|| auth_snapshot_allows_cross_format_candidate( && (matches_client_format
self.auth_snapshot, || auth_snapshot_allows_cross_format_candidate(
self.requested_model, self.auth_snapshot,
candidate, self.requested_model,
enable_model_directives, candidate,
) enable_model_directives,
))
} }
fn skipped_candidate_allowed( fn skipped_candidate_allowed(
@@ -118,13 +121,14 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
let enable_model_directives = self.model_directive_enabled_api_formats.contains( let enable_model_directives = self.model_directive_enabled_api_formats.contains(
&crate::ai_serving::normalize_api_format_alias(candidate_api_format), &crate::ai_serving::normalize_api_format_alias(candidate_api_format),
); );
matches_client_format routing_policy_allows_provider(self.routing_policy, &skipped_candidate.candidate)
|| auth_snapshot_allows_cross_format_candidate( && (matches_client_format
self.auth_snapshot, || auth_snapshot_allows_cross_format_candidate(
self.requested_model, self.auth_snapshot,
&skipped_candidate.candidate, self.requested_model,
enable_model_directives, &skipped_candidate.candidate,
) enable_model_directives,
))
} }
fn candidate_key(&self, candidate: &Self::Candidate) -> String { fn candidate_key(&self, candidate: &Self::Candidate) -> String {
@@ -144,6 +148,7 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
require_streaming: bool, require_streaming: bool,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
auth_snapshot: &GatewayAuthApiKeySnapshot, auth_snapshot: &GatewayAuthApiKeySnapshot,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
use_api_format_alias_match: bool, use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode, key_mode: LocalCandidatePreselectionKeyMode,
@@ -166,6 +171,7 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
require_streaming, require_streaming,
required_capabilities, required_capabilities,
auth_snapshot, auth_snapshot,
routing_policy,
client_session_affinity, client_session_affinity,
use_api_format_alias_match, use_api_format_alias_match,
key_mode, key_mode,
@@ -182,6 +188,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
require_streaming: bool, require_streaming: bool,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
auth_snapshot: &GatewayAuthApiKeySnapshot, auth_snapshot: &GatewayAuthApiKeySnapshot,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
use_api_format_alias_match: bool, use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode, key_mode: LocalCandidatePreselectionKeyMode,
@@ -213,6 +220,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
require_streaming, require_streaming,
required_capabilities, required_capabilities,
auth_snapshot, auth_snapshot,
routing_policy,
client_session_affinity, client_session_affinity,
use_api_format_alias_match, use_api_format_alias_match,
key_mode, key_mode,
@@ -230,6 +238,7 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
require_streaming: bool, require_streaming: bool,
required_capabilities: Option<serde_json::Value>, required_capabilities: Option<serde_json::Value>,
auth_snapshot: GatewayAuthApiKeySnapshot, auth_snapshot: GatewayAuthApiKeySnapshot,
routing_policy: Option<ResolvedRoutingPolicy>,
client_session_affinity: Option<ClientSessionAffinity>, client_session_affinity: Option<ClientSessionAffinity>,
use_api_format_alias_match: bool, use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode, key_mode: LocalCandidatePreselectionKeyMode,
@@ -253,6 +262,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
require_streaming: bool, require_streaming: bool,
required_capabilities: Option<&serde_json::Value>, required_capabilities: Option<&serde_json::Value>,
auth_snapshot: &GatewayAuthApiKeySnapshot, auth_snapshot: &GatewayAuthApiKeySnapshot,
routing_policy: Option<&ResolvedRoutingPolicy>,
client_session_affinity: Option<&ClientSessionAffinity>, client_session_affinity: Option<&ClientSessionAffinity>,
use_api_format_alias_match: bool, use_api_format_alias_match: bool,
key_mode: LocalCandidatePreselectionKeyMode, key_mode: LocalCandidatePreselectionKeyMode,
@@ -283,6 +293,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
require_streaming, require_streaming,
required_capabilities: required_capabilities.cloned(), required_capabilities: required_capabilities.cloned(),
auth_snapshot: auth_snapshot.clone(), auth_snapshot: auth_snapshot.clone(),
routing_policy: routing_policy.cloned(),
client_session_affinity: client_session_affinity.cloned(), client_session_affinity: client_session_affinity.cloned(),
use_api_format_alias_match, use_api_format_alias_match,
key_mode, key_mode,
@@ -620,16 +631,17 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
candidate_api_format: &str, candidate_api_format: &str,
enable_model_directives: bool, enable_model_directives: bool,
) -> bool { ) -> bool {
matches_client_api_format( routing_policy_allows_provider(self.routing_policy.as_ref(), candidate)
self.use_api_format_alias_match, && (matches_client_api_format(
candidate_api_format, self.use_api_format_alias_match,
&self.client_api_format, candidate_api_format,
) || auth_snapshot_allows_cross_format_candidate( &self.client_api_format,
&self.auth_snapshot, ) || auth_snapshot_allows_cross_format_candidate(
&self.requested_model, &self.auth_snapshot,
candidate, &self.requested_model,
enable_model_directives, candidate,
) enable_model_directives,
))
} }
fn skipped_candidate_allowed_for_page( fn skipped_candidate_allowed_for_page(
@@ -638,16 +650,17 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
candidate_api_format: &str, candidate_api_format: &str,
enable_model_directives: bool, enable_model_directives: bool,
) -> bool { ) -> bool {
matches_client_api_format( routing_policy_allows_provider(self.routing_policy.as_ref(), &skipped_candidate.candidate)
self.use_api_format_alias_match, && (matches_client_api_format(
candidate_api_format, self.use_api_format_alias_match,
&self.client_api_format, candidate_api_format,
) || auth_snapshot_allows_cross_format_candidate( &self.client_api_format,
&self.auth_snapshot, ) || auth_snapshot_allows_cross_format_candidate(
&self.requested_model, &self.auth_snapshot,
&skipped_candidate.candidate, &self.requested_model,
enable_model_directives, &skipped_candidate.candidate,
) enable_model_directives,
))
} }
} }
@@ -739,6 +752,18 @@ pub(crate) fn auth_snapshot_allows_cross_format_candidate(
true true
} }
fn routing_policy_allows_provider(
routing_policy: Option<&ResolvedRoutingPolicy>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> bool {
match routing_policy {
Some(policy) => policy
.ranking_overlay
.provider_allowed(candidate.provider_id.as_str()),
None => true,
}
}
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -887,6 +912,7 @@ mod tests {
None, None,
&auth_snapshot, &auth_snapshot,
None, None,
None,
true, true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
) )
@@ -943,6 +969,7 @@ mod tests {
None, None,
&auth_snapshot, &auth_snapshot,
None, None,
None,
true, true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
) )

View File

@@ -1,10 +1,27 @@
use std::collections::BTreeMap;
use aether_ai_serving::{run_ai_authenticated_decision_input, AiAuthenticatedDecisionInputPort}; use aether_ai_serving::{run_ai_authenticated_decision_input, AiAuthenticatedDecisionInputPort};
use aether_routing_core::{
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
RoutingCandidateTrace, RoutingDecisionTrace, RoutingPoolExpansionTrace, RoutingRulePhase,
};
use aether_scheduler_core::ClientSessionAffinity; use aether_scheduler_core::ClientSessionAffinity;
use async_trait::async_trait; use async_trait::async_trait;
use http::StatusCode;
use http::{HeaderMap, HeaderName, HeaderValue};
use serde_json::{json, Value};
use tracing::warn;
use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::{ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, PlannerAppState}; use crate::ai_serving::{ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::client_session_affinity::client_session_affinity_from_request;
use crate::clock::current_unix_secs; use crate::clock::current_unix_secs;
use crate::{AppState, GatewayError}; use crate::routing::{
apply_routing_mutation_plan, build_routing_trace_seed, resolve_gateway_routing_policy,
select_gateway_routing_group, GatewayRoutingPolicyInput, GatewayRoutingSelectionError,
GatewayRoutingSelectionInput, ROUTING_GROUP_HEADER,
};
use crate::{AiExecutionDecision, AppState, GatewayError};
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub(crate) struct ResolvedLocalDecisionAuthInput { pub(crate) struct ResolvedLocalDecisionAuthInput {
@@ -21,6 +38,9 @@ pub(crate) struct LocalRequestedModelDecisionInput {
pub(crate) required_capabilities: Option<serde_json::Value>, pub(crate) required_capabilities: Option<serde_json::Value>,
pub(crate) request_auth_channel: Option<String>, pub(crate) request_auth_channel: Option<String>,
pub(crate) client_session_affinity: Option<ClientSessionAffinity>, pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
} }
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -31,6 +51,92 @@ pub(crate) struct LocalAuthenticatedDecisionInput {
pub(crate) client_session_affinity: Option<ClientSessionAffinity>, pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
} }
#[derive(Debug, Clone)]
pub(crate) struct LocalRoutingRequestContext {
pub(crate) group_id: Option<String>,
pub(crate) group_version: Option<i64>,
pub(crate) group_config_json: Value,
pub(crate) selection_source: String,
pub(crate) client_api_format: String,
pub(crate) effective_body_json: Value,
pub(crate) effective_headers: HeaderMap,
}
impl LocalRequestedModelDecisionInput {
pub(crate) fn effective_body_json<'a>(&'a self, fallback: &'a Value) -> &'a Value {
self.routing_context
.as_ref()
.map(|context| &context.effective_body_json)
.unwrap_or(fallback)
}
pub(crate) fn effective_headers<'a>(&'a self, fallback: &'a HeaderMap) -> &'a HeaderMap {
self.routing_context
.as_ref()
.map(|context| &context.effective_headers)
.unwrap_or(fallback)
}
}
pub(crate) fn apply_provider_request_routing_policy_to_decision(
input: &LocalRequestedModelDecisionInput,
decision: &mut AiExecutionDecision,
) -> Result<(), GatewayError> {
let Some(context) = input.routing_context.as_ref() else {
return Ok(());
};
let provider_api_format = decision
.provider_api_format
.as_deref()
.unwrap_or(context.client_api_format.as_str());
let resolved_model = decision
.mapped_model
.as_deref()
.or(decision.model_name.as_deref())
.unwrap_or(input.requested_model.as_str());
let original_provider_request_body = decision.provider_request_body.clone();
let mut provider_request_body = original_provider_request_body
.clone()
.unwrap_or(serde_json::Value::Null);
let mut provider_headers = btree_headers_to_header_map(&decision.provider_request_headers)?;
let provider_headers_json = headers_to_routing_value(&provider_headers);
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
group_id: context.group_id.as_deref(),
group_version: context.group_version,
group_config_json: &context.group_config_json,
selection_source: context.selection_source.as_str(),
requested_model: input.requested_model.as_str(),
resolved_model,
api_format: provider_api_format,
user_id: Some(input.auth_context.user_id.as_str()),
api_key_id: Some(input.auth_context.api_key_id.as_str()),
headers: &provider_headers_json,
body: &provider_request_body,
phase: RoutingRulePhase::ProviderRequest,
})?;
ensure_report_context_routing_trace(input, decision, &policy);
if policy.mutation_plan.is_empty() {
return Ok(());
}
if original_provider_request_body.is_none() && !policy.mutation_plan.body_patch.is_empty() {
return Err(GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: "routing provider_request body patch cannot be applied to a binary or empty upstream body".to_string(),
});
}
apply_routing_mutation_plan(
&mut provider_request_body,
&mut provider_headers,
&policy.mutation_plan,
)?;
decision.provider_request_headers = header_map_to_btree_headers(&provider_headers);
if original_provider_request_body.is_some() {
decision.provider_request_body = Some(provider_request_body);
}
update_report_context_provider_request_mutation(decision, &policy);
Ok(())
}
struct GatewayAuthenticatedDecisionInputPort<'a> { struct GatewayAuthenticatedDecisionInputPort<'a> {
state: PlannerAppState<'a>, state: PlannerAppState<'a>,
now_unix_secs: u64, now_unix_secs: u64,
@@ -99,9 +205,154 @@ pub(crate) fn build_local_requested_model_decision_input(
required_capabilities: resolved_input.required_capabilities, required_capabilities: resolved_input.required_capabilities,
request_auth_channel: None, request_auth_channel: None,
client_session_affinity: None, client_session_affinity: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
} }
} }
pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
state: &AppState,
parts: &http::request::Parts,
input: &mut LocalRequestedModelDecisionInput,
body_json: &Value,
client_api_format: &str,
) -> Result<(), GatewayError> {
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
let selected_group = match state.routing_group_read_repository() {
Some(repository) => {
let user_group_ids = match state
.list_user_groups_for_user(&input.auth_context.user_id)
.await
{
Ok(groups) => groups.into_iter().map(|group| group.id).collect::<Vec<_>>(),
Err(error) => {
warn!(
user_id = %input.auth_context.user_id,
error = ?error,
"gateway routing profile user group lookup failed"
);
Vec::new()
}
};
let selection = select_gateway_routing_group(
repository.as_ref(),
GatewayRoutingSelectionInput {
explicit_group: explicit_group.as_deref(),
user_id: Some(input.auth_context.user_id.as_str()),
api_key_id: Some(input.auth_context.api_key_id.as_str()),
user_group_ids: &user_group_ids,
},
)
.await
.map_err(routing_selection_error)?;
selection.group.map(|group| {
(
Some(group.id),
Some(group.version),
group.config_json,
selection.source,
)
})
}
None => {
if explicit_group
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty())
{
return Err(routing_selection_error(
GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()),
));
}
None
}
};
let Some((group_id, group_version, group_config_json, selection_source)) = selected_group
else {
input.client_session_affinity =
client_session_affinity_from_request(&parts.headers, Some(body_json));
input.routing_policy = None;
input.routing_trace_seed = None;
input.routing_context = None;
return Ok(());
};
let headers_json = headers_to_routing_value(&parts.headers);
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
group_id: group_id.as_deref(),
group_version,
group_config_json: &group_config_json,
selection_source: selection_source.as_str(),
requested_model: input.requested_model.as_str(),
resolved_model: input.requested_model.as_str(),
api_format: client_api_format,
user_id: Some(input.auth_context.user_id.as_str()),
api_key_id: Some(input.auth_context.api_key_id.as_str()),
headers: &headers_json,
body: body_json,
phase: RoutingRulePhase::ClientRequest,
})?;
let mut effective_body_json = body_json.clone();
let mut effective_headers = parts.headers.clone();
apply_routing_mutation_plan(
&mut effective_body_json,
&mut effective_headers,
&policy.mutation_plan,
)?;
let mut requested_model_changed = false;
if let Some(mut mutated_model) = extract_standard_requested_model(&effective_body_json) {
mutated_model = mutated_model.trim().to_string();
if !mutated_model.is_empty() && mutated_model != input.requested_model {
input.requested_model = mutated_model;
requested_model_changed = true;
}
}
if requested_model_changed {
input.required_capabilities = PlannerAppState::new(state)
.resolve_request_candidate_required_capabilities(
&input.auth_context.user_id,
&input.auth_context.api_key_id,
Some(input.requested_model.as_str()),
input.required_capabilities.as_ref(),
)
.await;
}
let effective_headers_json = headers_to_routing_value(&effective_headers);
input.client_session_affinity =
client_session_affinity_from_request(&effective_headers, Some(&effective_body_json));
let mut final_policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
group_id: group_id.as_deref(),
group_version,
group_config_json: &group_config_json,
selection_source: selection_source.as_str(),
requested_model: input.requested_model.as_str(),
resolved_model: input.requested_model.as_str(),
api_format: client_api_format,
user_id: Some(input.auth_context.user_id.as_str()),
api_key_id: Some(input.auth_context.api_key_id.as_str()),
headers: &effective_headers_json,
body: &effective_body_json,
phase: RoutingRulePhase::ClientRequest,
})?;
final_policy.mutation_plan = policy.mutation_plan.clone();
input.routing_trace_seed = Some(build_routing_trace_seed(&final_policy, client_api_format));
input.routing_policy = Some(final_policy);
input.routing_context = Some(LocalRoutingRequestContext {
group_id,
group_version,
group_config_json,
selection_source,
client_api_format: client_api_format.to_string(),
effective_body_json,
effective_headers,
});
Ok(())
}
pub(crate) fn build_local_authenticated_decision_input( pub(crate) fn build_local_authenticated_decision_input(
resolved_input: ResolvedLocalDecisionAuthInput, resolved_input: ResolvedLocalDecisionAuthInput,
) -> LocalAuthenticatedDecisionInput { ) -> LocalAuthenticatedDecisionInput {
@@ -132,3 +383,506 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
) )
.await .await
} }
fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError {
GatewayError::Client {
status: StatusCode::FORBIDDEN,
message: error.to_string(),
}
}
fn headers_to_routing_value(headers: &http::HeaderMap) -> Value {
let mut object = serde_json::Map::new();
for (name, value) in headers {
if let Ok(value) = value.to_str() {
object.insert(name.as_str().to_ascii_lowercase(), json!(value));
}
}
Value::Object(object)
}
fn routing_header_value_str(headers: &http::HeaderMap, key: &str) -> Option<String> {
headers
.get(key)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
fn btree_headers_to_header_map(
headers: &BTreeMap<String, String>,
) -> Result<HeaderMap, GatewayError> {
let mut output = HeaderMap::new();
for (name, value) in headers {
let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: format!("invalid provider request header name in routing mutation: {err}"),
})?;
let value = HeaderValue::from_str(value).map_err(|err| GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: format!("invalid provider request header value in routing mutation: {err}"),
})?;
output.insert(name, value);
}
Ok(output)
}
fn header_map_to_btree_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
headers
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.as_str().to_string(), value.to_string()))
})
.collect()
}
fn update_report_context_provider_request_mutation(
decision: &mut AiExecutionDecision,
policy: &ResolvedRoutingPolicy,
) {
let Some(serde_json::Value::Object(object)) = decision.report_context.as_mut() else {
return;
};
let body_paths = policy
.mutation_plan
.body_patch
.iter()
.map(|operation| operation.path().to_string())
.collect::<Vec<_>>();
let header_names = policy
.mutation_plan
.header_patch
.iter()
.map(|operation| operation.name().to_string())
.collect::<Vec<_>>();
let trace_patch_summary = serde_json::json!({
"body_paths": body_paths,
"header_names": header_names,
});
if let Some(serde_json::Value::Object(routing_trace)) = object.get_mut("routing_trace") {
routing_trace.insert(
"provider_request_patch_summary".to_string(),
trace_patch_summary.clone(),
);
}
object.insert(
"provider_request_headers".to_string(),
serde_json::json!(decision.provider_request_headers),
);
object.insert(
"routing_provider_request_patch_summary".to_string(),
serde_json::json!({
"body_paths": trace_patch_summary["body_paths"].clone(),
"header_names": trace_patch_summary["header_names"].clone(),
"matched_rules": policy
.matched_rules
.iter()
.map(|rule| rule.id.clone())
.collect::<Vec<_>>()
}),
);
}
fn ensure_report_context_routing_trace(
input: &LocalRequestedModelDecisionInput,
decision: &mut AiExecutionDecision,
policy: &ResolvedRoutingPolicy,
) {
let Some(serde_json::Value::Object(object)) = decision.report_context.as_mut() else {
return;
};
if object.get("routing_trace").is_some() {
return;
}
let client_api_format = decision
.client_api_format
.as_deref()
.or_else(|| {
input
.routing_context
.as_ref()
.map(|context| context.client_api_format.as_str())
})
.unwrap_or_default();
let mut trace = input
.routing_trace_seed
.clone()
.unwrap_or_else(|| build_routing_trace_seed(policy, client_api_format));
let candidate_group_id = object
.get("candidate_group_id")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
let pool_key_index = object
.get("pool_key_index")
.and_then(Value::as_u64)
.and_then(|value| u32::try_from(value).ok());
let is_pool_expansion = candidate_group_id.is_some() && pool_key_index.is_some();
let candidate_kind = if is_pool_expansion {
CandidateKind::PoolGroup
} else {
CandidateKind::Provider
};
let provider_id = candidate_group_id
.clone()
.or_else(|| decision.provider_id.clone())
.unwrap_or_default();
let endpoint_id = decision.endpoint_id.clone().unwrap_or_default();
let model_id = object
.get("model_id")
.and_then(Value::as_str)
.map(ToOwned::to_owned)
.or_else(|| decision.mapped_model.clone())
.or_else(|| decision.model_name.clone())
.unwrap_or_else(|| input.requested_model.clone());
let key_id = decision.key_id.clone().filter(|_| !is_pool_expansion);
let provider_priority = object
.get("provider_priority")
.and_then(Value::as_i64)
.and_then(|value| i32::try_from(value).ok())
.unwrap_or_default();
let key_priority = object
.get("priority_slot")
.and_then(Value::as_i64)
.and_then(|value| i32::try_from(value).ok())
.unwrap_or_default();
trace.global_candidates.push(RoutingCandidateTrace {
candidate_kind,
provider_id: provider_id.clone(),
endpoint_id,
model_id: model_id.clone(),
key_id: key_id.clone(),
ranking_vector: rank_vector_for_candidate(
&policy.ranking_overlay,
&RoutingCandidateFacts {
candidate_kind,
provider_id: provider_id.clone(),
endpoint_id: decision.endpoint_id.clone().unwrap_or_default(),
model_id,
key_id,
provider_priority,
key_priority,
},
),
skip_reason: None,
selected_order: object
.get("candidate_index")
.and_then(Value::as_u64)
.and_then(|value| u32::try_from(value).ok()),
});
if is_pool_expansion {
if let (Some(pool_group_id), Some(key_id)) = (candidate_group_id, decision.key_id.clone()) {
trace.pool_expansion.push(RoutingPoolExpansionTrace {
pool_group_id,
key_id,
pool_ranking_vector: Vec::new(),
pool_skip_reason: None,
selected_order: pool_key_index,
});
}
}
object.insert("routing_trace".to_string(), serde_json::json!(trace));
}
#[cfg(test)]
mod tests {
use super::*;
fn sample_auth_context() -> ExecutionRuntimeAuthContext {
ExecutionRuntimeAuthContext {
user_id: "user-1".to_string(),
api_key_id: "api-key-1".to_string(),
username: None,
api_key_name: None,
balance_remaining: None,
access_allowed: true,
api_key_is_standalone: false,
}
}
fn sample_auth_snapshot() -> GatewayAuthApiKeySnapshot {
GatewayAuthApiKeySnapshot {
user_id: "user-1".to_string(),
username: "alice".to_string(),
email: None,
user_role: "user".to_string(),
user_auth_source: "local".to_string(),
user_is_active: true,
user_is_deleted: false,
user_rate_limit: None,
user_allowed_providers: None,
user_allowed_api_formats: None,
user_allowed_models: None,
api_key_id: "api-key-1".to_string(),
api_key_name: Some("default".to_string()),
api_key_is_active: true,
api_key_is_locked: false,
api_key_is_standalone: false,
api_key_rate_limit: None,
api_key_concurrent_limit: None,
api_key_expires_at_unix_secs: None,
api_key_allowed_providers: None,
api_key_allowed_api_formats: None,
api_key_allowed_models: None,
currently_usable: true,
}
}
fn sample_decision_input() -> LocalRequestedModelDecisionInput {
LocalRequestedModelDecisionInput {
auth_context: sample_auth_context(),
requested_model: "gpt-5".to_string(),
auth_snapshot: sample_auth_snapshot(),
required_capabilities: None,
request_auth_channel: None,
client_session_affinity: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: Some(LocalRoutingRequestContext {
group_id: Some("group-1".to_string()),
group_version: Some(3),
selection_source: "explicit_header".to_string(),
client_api_format: "openai:chat".to_string(),
effective_body_json: json!({"model":"gpt-5"}),
effective_headers: HeaderMap::new(),
group_config_json: json!({
"allowed_models": ["gpt-5"],
"rules": [{
"id": "provider-patch",
"priority": 1,
"enabled": true,
"phase": "provider_request",
"conditions": {},
"actions": [
{
"type": "json_patch_body",
"patch": [{
"op": "add",
"path": "/metadata/routing",
"value": "provider"
}]
},
{
"type": "patch_headers",
"patch": [{
"op": "set",
"name": "x-provider-route",
"value": "provider"
}]
}
]
}]
}),
}),
}
}
fn sample_decision() -> AiExecutionDecision {
AiExecutionDecision {
action: "execution_runtime_sync_decision".to_string(),
decision_kind: Some("openai_chat_sync".to_string()),
execution_strategy: None,
conversion_mode: None,
request_id: Some("trace-1".to_string()),
candidate_id: Some("candidate-1".to_string()),
provider_name: Some("provider".to_string()),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("key-1".to_string()),
upstream_base_url: None,
upstream_url: None,
provider_request_method: None,
auth_header: None,
auth_value: None,
provider_api_format: Some("openai:chat".to_string()),
client_api_format: Some("openai:chat".to_string()),
provider_contract: None,
client_contract: None,
model_name: Some("gpt-5".to_string()),
mapped_model: Some("gpt-5".to_string()),
prompt_cache_key: None,
extra_headers: BTreeMap::new(),
provider_request_headers: BTreeMap::from([(
"content-type".to_string(),
"application/json".to_string(),
)]),
provider_request_body: Some(json!({"model":"gpt-5","metadata":{}})),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
proxy: None,
transport_profile: None,
timeouts: None,
upstream_is_stream: false,
report_kind: Some("local_sync_success".to_string()),
report_context: Some(json!({
"candidate_index": 0,
"retry_index": 0,
"model_id": "model-1"
})),
auth_context: Some(sample_auth_context()),
}
}
fn set_provider_request_rules(input: &mut LocalRequestedModelDecisionInput, actions: Value) {
let config = json!({
"allowed_models": ["gpt-5"],
"rules": [{
"id": "provider-patch",
"priority": 1,
"enabled": true,
"phase": "provider_request",
"conditions": {},
"actions": actions
}]
});
input
.routing_context
.as_mut()
.expect("sample input should include routing context")
.group_config_json = config;
}
#[test]
fn provider_request_routing_policy_mutates_decision_body_headers_and_report_context() {
let input = sample_decision_input();
let mut decision = sample_decision();
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
.expect("provider routing mutation should apply");
assert_eq!(
decision.provider_request_body.as_ref().unwrap()["metadata"]["routing"],
json!("provider")
);
assert_eq!(
decision
.provider_request_headers
.get("x-provider-route")
.map(String::as_str),
Some("provider")
);
let report_context = decision.report_context.as_ref().unwrap();
assert_eq!(
report_context["routing_provider_request_patch_summary"]["matched_rules"],
json!(["provider-patch"])
);
assert_eq!(
report_context["routing_trace"]["provider_request_patch_summary"]["body_paths"],
json!(["/metadata/routing"])
);
assert_eq!(
report_context["routing_trace"]["global_candidates"][0]["provider_id"],
json!("provider-1")
);
}
#[test]
fn provider_request_routing_policy_rejects_body_patch_without_json_body() {
let input = sample_decision_input();
let mut decision = sample_decision();
decision.provider_request_body = None;
decision.provider_request_body_base64 = Some("AA==".to_string());
let error = apply_provider_request_routing_policy_to_decision(&input, &mut decision)
.expect_err("provider body patch should reject binary upstream bodies");
match error {
GatewayError::Client { status, message } => {
assert_eq!(status, StatusCode::BAD_REQUEST);
assert!(message.contains("binary or empty upstream body"));
}
other => panic!("unexpected error: {other:?}"),
}
assert!(
decision
.report_context
.as_ref()
.and_then(|context| context.get("routing_trace"))
.is_some(),
"failed provider_request mutation should still seed routing trace"
);
}
#[test]
fn provider_request_routing_policy_allows_header_patch_without_json_body() {
let mut input = sample_decision_input();
set_provider_request_rules(
&mut input,
json!([{
"type": "patch_headers",
"patch": [{
"op": "set",
"name": "x-provider-route",
"value": "header-only"
}]
}]),
);
let mut decision = sample_decision();
decision.provider_request_body = None;
decision.provider_request_body_base64 = Some("AA==".to_string());
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
.expect("header-only provider routing mutation should apply without JSON body");
assert_eq!(decision.provider_request_body, None);
assert_eq!(
decision
.provider_request_headers
.get("x-provider-route")
.map(String::as_str),
Some("header-only")
);
assert_eq!(
decision.report_context.as_ref().unwrap()["routing_trace"]
["provider_request_patch_summary"]["header_names"],
json!(["x-provider-route"])
);
}
#[test]
fn provider_request_routing_trace_records_pool_expansion_candidate() {
let input = sample_decision_input();
let mut decision = sample_decision();
decision.report_context = Some(json!({
"candidate_index": 2,
"retry_index": 2,
"model_id": "model-1",
"candidate_group_id": "pool-group-1",
"pool_key_index": 1,
"provider_priority": 7,
"priority_slot": 3
}));
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
.expect("provider routing mutation should seed pool trace");
let routing_trace = &decision.report_context.as_ref().unwrap()["routing_trace"];
assert_eq!(
routing_trace["global_candidates"][0]["candidate_kind"],
json!("pool_group")
);
assert_eq!(
routing_trace["global_candidates"][0]["provider_id"],
json!("pool-group-1")
);
assert_eq!(routing_trace["global_candidates"][0]["key_id"], Value::Null);
assert_eq!(
routing_trace["pool_expansion"][0]["pool_group_id"],
json!("pool-group-1")
);
assert_eq!(routing_trace["pool_expansion"][0]["key_id"], json!("key-1"));
assert_eq!(
routing_trace["pool_expansion"][0]["selected_order"],
json!(1)
);
}
}

View File

@@ -33,7 +33,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -55,6 +55,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
) )
@@ -70,7 +71,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
maybe_build_local_same_format_provider_decision_payload_for_candidate( maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -100,7 +101,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -122,6 +123,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
) )
@@ -137,7 +139,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
maybe_build_local_same_format_provider_decision_payload_for_candidate( maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }

View File

@@ -13,6 +13,7 @@ use crate::ai_serving::planner::candidate_metadata::{
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate; use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::planner::common::extract_requested_model_from_request; use crate::ai_serving::planner::common::extract_requested_model_from_request;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
@@ -39,19 +40,21 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
body_json: &serde_json::Value, body_json: &serde_json::Value,
spec: LocalSameFormatProviderSpec, spec: LocalSameFormatProviderSpec,
) -> Option<LocalSameFormatProviderDecisionInput> { ) -> Result<Option<LocalSameFormatProviderDecisionInput>, GatewayError> {
let spec_metadata = local_same_format_provider_spec_metadata(spec); let spec_metadata = local_same_format_provider_spec_metadata(spec);
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
return None; return Ok(None);
}; };
let requested_model = extract_requested_model_from_request( let Some(requested_model) = extract_requested_model_from_request(
parts, parts,
body_json, body_json,
spec_metadata spec_metadata
.requested_model_family .requested_model_family
.expect("same-format provider specs should declare requested-model family"), .expect("same-format provider specs should declare requested-model family"),
)?; ) else {
return Ok(None);
};
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
state, state,
@@ -62,7 +65,7 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
.await .await
{ {
Ok(Some(resolved_input)) => resolved_input, Ok(Some(resolved_input)) => resolved_input,
Ok(None) => return None, Ok(None) => return Ok(None),
Err(err) => { Err(err) => {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
@@ -70,14 +73,31 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
error = ?err, error = ?err,
"gateway local same-format decision auth snapshot read failed" "gateway local same-format decision auth snapshot read failed"
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
spec_metadata.api_format,
)
.await
{
warn!(
trace_id = %trace_id,
api_format = spec_metadata.api_format,
error = ?err,
"gateway local same-format decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }
pub(crate) async fn materialize_local_same_format_provider_candidate_attempts( pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
@@ -114,6 +134,7 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -211,6 +232,7 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,

View File

@@ -6,6 +6,7 @@ use crate::ai_serving::planner::candidate_materialization::{
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data, mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
mark_skipped_local_execution_candidate_with_failure_diagnostic, mark_skipped_local_execution_candidate_with_failure_diagnostic,
}; };
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind, build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
}; };
@@ -22,7 +23,7 @@ use crate::ai_serving::transport::{
}; };
use crate::{ use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value, append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, AiExecutionDecision, AppState, GatewayError,
}; };
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate; use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
@@ -40,7 +41,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
input: &LocalSameFormatProviderDecisionInput, input: &LocalSameFormatProviderDecisionInput,
attempt: LocalSameFormatProviderCandidateAttempt, attempt: LocalSameFormatProviderCandidateAttempt,
spec: LocalSameFormatProviderSpec, spec: LocalSameFormatProviderSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_same_format_provider_spec_metadata(spec); let spec_metadata = local_same_format_provider_spec_metadata(spec);
let LocalSameFormatProviderCandidateAttempt { let LocalSameFormatProviderCandidateAttempt {
eligible, eligible,
@@ -51,10 +52,13 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let (execution_strategy, conversion_mode) = let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(spec_metadata.api_format, spec_metadata.api_format); ai_local_execution_contract_for_formats(spec_metadata.api_format, spec_metadata.api_format);
let resolved = resolve_local_same_format_provider_candidate_payload_parts( let Some(resolved) = resolve_local_same_format_provider_candidate_payload_parts(
state, parts, trace_id, body_json, input, &attempt, spec, state, parts, trace_id, body_json, input, &attempt, spec,
) )
.await?; .await
else {
return Ok(None);
};
let prompt_cache_key = resolved let prompt_cache_key = resolved
.provider_request_body .provider_request_body
@@ -66,7 +70,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
let proxy = state let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport) .resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await; .await;
let transport_profile = resolve_transport_profile(&resolved.transport); let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let mut extra_fields = serde_json::Map::new(); let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) = if let Some(proxy_value) =
build_request_trace_proxy_value(Some(&resolved.transport), proxy.as_ref()) build_request_trace_proxy_value(Some(&resolved.transport), proxy.as_ref())
@@ -85,6 +92,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
); );
} }
let provider_api_format = resolved.provider_api_format.clone(); let provider_api_format = resolved.provider_api_format.clone();
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_local_failover_policy_to_value( let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value( append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts { build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -112,7 +120,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
body_rules: resolved.transport.endpoint.body_rules.as_ref(), body_rules: resolved.transport.endpoint.body_rules.as_ref(),
provider_request_method: Some(serde_json::Value::Null), provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers), provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -149,43 +157,44 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
upstream_url, upstream_url,
provider_request_headers, provider_request_headers,
provider_request_body, provider_request_body,
transport_profile: _,
} = resolved; } = resolved;
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: spec_metadata.require_streaming,
decision_is_stream: spec_metadata.require_streaming, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.to_string(),
candidate_id: candidate_id.to_string(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: None,
provider_request_method: None, auth_header,
auth_header, auth_value,
auth_value, provider_api_format,
provider_api_format, client_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(), model_name: input.requested_model.clone(),
model_name: input.requested_model.clone(), mapped_model,
mapped_model, prompt_cache_key,
prompt_cache_key, provider_request_headers,
provider_request_headers, provider_request_body: Some(provider_request_body),
provider_request_body: Some(provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: Some("application/json".to_string()),
content_type: Some("application/json".to_string()), proxy,
proxy, transport_profile,
transport_profile, timeouts: resolve_transport_execution_timeouts(&transport),
timeouts: resolve_transport_execution_timeouts(&transport), upstream_is_stream,
upstream_is_stream, report_kind: Some(report_kind.to_string()),
report_kind: Some(report_kind.to_string()), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
)) Ok(Some(decision))
} }
pub(super) async fn mark_skipped_local_same_format_provider_candidate( pub(super) async fn mark_skipped_local_same_format_provider_candidate(

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value; use serde_json::Value;
use crate::ai_serving::planner::common::{ use crate::ai_serving::planner::common::{
@@ -12,7 +13,8 @@ use crate::ai_serving::transport::antigravity::{
AntigravityRequestEnvelopeSupport, AntigravityRequestSideSupport, AntigravityRequestEnvelopeSupport, AntigravityRequestSideSupport,
}; };
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_same_format_provider_headers, SameFormatProviderHeadersInput, build_grok_browser_headers, build_grok_upstream_url, build_same_format_provider_headers,
GrokHeaderInput, SameFormatProviderHeadersInput, GROK_CHAT_PATH,
}; };
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot}; use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
use crate::AppState; use crate::AppState;
@@ -96,6 +98,7 @@ pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
pub(super) upstream_url: String, pub(super) upstream_url: String,
pub(super) provider_request_headers: BTreeMap<String, String>, pub(super) provider_request_headers: BTreeMap<String, String>,
pub(super) provider_request_body: Value, pub(super) provider_request_body: Value,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
} }
pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts( pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
@@ -125,6 +128,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
Some(&input.requested_model), Some(&input.requested_model),
) )
.await; .await;
let effective_headers = input.effective_headers(&parts.headers);
let Some(mut base_provider_request_body) = let Some(mut base_provider_request_body) =
super::super::request::build_same_format_provider_request_body( super::super::request::build_same_format_provider_request_body(
@@ -133,7 +137,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
&prepared.mapped_model, &prepared.mapped_model,
spec, spec,
prepared.transport.endpoint.body_rules.as_ref(), prepared.transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
prepared.upstream_is_stream, prepared.upstream_is_stream,
prepared.force_body_stream_field, prepared.force_body_stream_field,
prepared.kiro_auth.as_ref(), prepared.kiro_auth.as_ref(),
@@ -246,16 +250,28 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
base_provider_request_body base_provider_request_body
}; };
let Some(upstream_url) = super::super::request::build_same_format_upstream_url( let is_grok = prepared
parts, .transport
&prepared.transport, .provider
&prepared.mapped_model, .provider_type
prepared.provider_api_format.as_str(), .trim()
spec, .eq_ignore_ascii_case("grok");
prepared.upstream_is_stream, let transport_profile =
prepared.kiro_auth.as_ref(), crate::ai_serving::transport::resolve_transport_profile(&prepared.transport);
Some(&provider_request_body), let Some(upstream_url) = (if is_grok {
) else { Some(build_grok_upstream_url(&prepared.transport, GROK_CHAT_PATH))
} else {
super::super::request::build_same_format_upstream_url(
parts,
&prepared.transport,
&prepared.mapped_model,
prepared.provider_api_format.as_str(),
spec,
prepared.upstream_is_stream,
prepared.kiro_auth.as_ref(),
Some(&provider_request_body),
)
}) else {
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic( mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
state, state,
input, input,
@@ -278,9 +294,20 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
.as_ref() .as_ref()
.map(build_antigravity_static_identity_headers) .map(build_antigravity_static_identity_headers)
.unwrap_or_default(); .unwrap_or_default();
let Some(provider_request_headers) = let Some(provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput {
transport: &prepared.transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: prepared.transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
})
} else {
build_same_format_provider_headers(SameFormatProviderHeadersInput { build_same_format_provider_headers(SameFormatProviderHeadersInput {
headers: &parts.headers, headers: effective_headers,
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: body_json, original_request_body: body_json,
header_rules: prepared.transport.endpoint.header_rules.as_ref(), header_rules: prepared.transport.endpoint.header_rules.as_ref(),
@@ -295,7 +322,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
.as_ref() .as_ref()
.map(|auth| auth.machine_id.as_str()), .map(|auth| auth.machine_id.as_str()),
}) })
else { }) else {
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic( mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
state, state,
input, input,
@@ -327,5 +354,6 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
upstream_url, upstream_url,
provider_request_headers, provider_request_headers,
provider_request_body, provider_request_body,
transport_profile,
}) })
} }

View File

@@ -31,7 +31,7 @@ pub(crate) struct LocalSameFormatProviderSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalSameFormatProviderDecisionInput, input: LocalSameFormatProviderDecisionInput,
spec: LocalSameFormatProviderSpec, spec: LocalSameFormatProviderSpec,
requested_model_family: RequestedModelFamily, requested_model_family: RequestedModelFamily,
@@ -42,7 +42,7 @@ pub(crate) struct LocalSameFormatProviderStreamAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalSameFormatProviderDecisionInput, input: LocalSameFormatProviderDecisionInput,
spec: LocalSameFormatProviderSpec, spec: LocalSameFormatProviderSpec,
requested_model_family: RequestedModelFamily, requested_model_family: RequestedModelFamily,
@@ -64,7 +64,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -85,8 +85,13 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state,
trace_id,
&input,
&effective_body_json,
spec,
) )
.await?; .await?;
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal( apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
@@ -103,7 +108,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
requested_model_family, requested_model_family,
@@ -128,7 +133,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -149,8 +154,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state,
trace_id,
&input,
&effective_body_json,
spec,
) )
.await?; .await?;
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal( apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
@@ -167,7 +177,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
requested_model_family, requested_model_family,
@@ -244,12 +254,12 @@ impl LocalSameFormatProviderSyncAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -257,7 +267,7 @@ impl LocalSameFormatProviderSyncAttemptSource<'_> {
match build_sync_plan_from_requested_model_family( match build_sync_plan_from_requested_model_family(
self.requested_model_family, self.requested_model_family,
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
) { ) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
@@ -282,12 +292,12 @@ impl LocalSameFormatProviderStreamAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -295,7 +305,7 @@ impl LocalSameFormatProviderStreamAttemptSource<'_> {
match build_stream_plan_from_requested_model_family( match build_stream_plan_from_requested_model_family(
self.requested_model_family, self.requested_model_family,
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
) { ) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
@@ -326,7 +336,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -347,6 +357,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
) )
@@ -365,7 +376,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate( let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };
@@ -411,7 +422,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
let Some(input) = resolve_local_same_format_provider_decision_input( let Some(input) = resolve_local_same_format_provider_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -432,6 +443,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source( let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
) )
@@ -450,7 +462,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate( let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -82,6 +82,18 @@ pub(crate) fn build_local_execution_report_context(
.client_session_affinity .client_session_affinity
.and_then(client_session_affinity_report_context_value) .and_then(client_session_affinity_report_context_value)
{ {
if let Some(client_family) = value
.as_object()
.and_then(|object| object.get("client_family"))
.and_then(Value::as_str)
.map(str::trim)
.filter(|client_family| !client_family.is_empty())
{
extra_fields.insert(
"client_family".to_string(),
Value::String(client_family.to_ascii_lowercase()),
);
}
extra_fields.insert( extra_fields.insert(
CLIENT_SESSION_AFFINITY_REPORT_CONTEXT_FIELD.to_string(), CLIENT_SESSION_AFFINITY_REPORT_CONTEXT_FIELD.to_string(),
value, value,

View File

@@ -29,7 +29,7 @@ use self::support::{
pub(crate) struct LocalGeminiFilesSyncAttemptSource<'a> { pub(crate) struct LocalGeminiFilesSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
body_base64: Option<&'a str>, body_base64: Option<&'a str>,
body_is_empty: bool, body_is_empty: bool,
trace_id: &'a str, trace_id: &'a str,
@@ -110,10 +110,11 @@ pub(crate) async fn build_local_gemini_files_sync_attempt_source_for_kind<'a>(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = let (candidates, candidate_count) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?; build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
if candidate_count == 0 { if candidate_count == 0 {
@@ -124,7 +125,7 @@ pub(crate) async fn build_local_gemini_files_sync_attempt_source_for_kind<'a>(
LocalGeminiFilesSyncAttemptSource { LocalGeminiFilesSyncAttemptSource {
state, state,
parts, parts,
body_json, body_json: effective_body_json,
body_base64, body_base64,
body_is_empty, body_is_empty,
trace_id, trace_id,
@@ -148,7 +149,7 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
}; };
let Some(input) = let Some(input) =
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -226,7 +227,7 @@ impl LocalGeminiFilesSyncAttemptSource<'_> {
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate( let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
self.state, self.state,
self.parts, self.parts,
self.body_json, &self.body_json,
self.body_base64, self.body_base64,
self.body_is_empty, self.body_is_empty,
self.trace_id, self.trace_id,
@@ -234,7 +235,7 @@ impl LocalGeminiFilesSyncAttemptSource<'_> {
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -272,7 +273,7 @@ impl LocalGeminiFilesStreamAttemptSource<'_> {
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -313,10 +314,11 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?; build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
@@ -333,7 +335,7 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
attempt, attempt,
spec, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -354,7 +356,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
}; };
let Some(input) = let Some(input) =
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -375,7 +377,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
attempt, attempt,
spec, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -402,10 +404,11 @@ async fn build_local_sync_plan_and_reports(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = let (mut source, _) =
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?; build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
@@ -423,7 +426,7 @@ async fn build_local_sync_plan_and_reports(
attempt, attempt,
spec, spec,
) )
.await .await?
else { else {
continue; continue;
}; };
@@ -454,7 +457,7 @@ async fn build_local_stream_plan_and_reports(
) -> Result<Vec<AiStreamAttempt>, GatewayError> { ) -> Result<Vec<AiStreamAttempt>, GatewayError> {
let spec_metadata = local_gemini_files_spec_metadata(spec); let spec_metadata = local_gemini_files_spec_metadata(spec);
let Some(input) = let Some(input) =
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
@@ -476,7 +479,7 @@ async fn build_local_stream_plan_and_reports(
attempt, attempt,
spec, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -1,6 +1,7 @@
use serde_json::json; use serde_json::json;
use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::report_context::{ use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, LocalExecutionReportContextParts, build_local_execution_report_context, LocalExecutionReportContextParts,
}; };
@@ -12,7 +13,7 @@ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_execution_timeouts, resolve_transport_profile,
}; };
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState}; use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{AiExecutionDecision, AppState}; use crate::{AiExecutionDecision, AppState, GatewayError};
use super::request::resolve_local_gemini_files_candidate_payload_parts; use super::request::resolve_local_gemini_files_candidate_payload_parts;
use super::support::{ use super::support::{
@@ -31,7 +32,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
input: &LocalGeminiFilesDecisionInput, input: &LocalGeminiFilesDecisionInput,
attempt: LocalGeminiFilesCandidateAttempt, attempt: LocalGeminiFilesCandidateAttempt,
spec: LocalGeminiFilesSpec, spec: LocalGeminiFilesSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_gemini_files_spec_metadata(spec); let spec_metadata = local_gemini_files_spec_metadata(spec);
let planner_state = PlannerAppState::new(state); let planner_state = PlannerAppState::new(state);
let attempt_identity = attempt.attempt_identity(); let attempt_identity = attempt.attempt_identity();
@@ -46,7 +47,10 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
&attempt, &attempt,
spec, spec,
) )
.await?; .await;
let Some(resolved) = resolved else {
return Ok(None);
};
let LocalGeminiFilesCandidateAttempt { let LocalGeminiFilesCandidateAttempt {
eligible, eligible,
candidate_id, candidate_id,
@@ -69,6 +73,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
} }
extra_fields.insert("file_key_id".to_string(), json!(candidate.key_id)); extra_fields.insert("file_key_id".to_string(), json!(candidate.key_id));
extra_fields.insert("file_name".to_string(), json!(resolved.file_name)); extra_fields.insert("file_name".to_string(), json!(resolved.file_name));
let effective_headers = input.effective_headers(&parts.headers);
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context, auth_context: &input.auth_context,
request_id: trace_id, request_id: trace_id,
@@ -94,7 +99,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
body_rules: transport.endpoint.body_rules.as_ref(), body_rules: transport.endpoint.body_rules.as_ref(),
provider_request_method: None, provider_request_method: None,
provider_request_headers: None, provider_request_headers: None,
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -119,45 +124,44 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
file_name: _, file_name: _,
} = resolved; } = resolved;
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: spec_metadata.require_streaming,
decision_is_stream: spec_metadata.require_streaming, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.clone(),
candidate_id: candidate_id.clone(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: Some(parts.method.to_string()),
provider_request_method: Some(parts.method.to_string()), auth_header: Some(auth_header),
auth_header: Some(auth_header), auth_value: Some(auth_value),
auth_value: Some(auth_value), provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(), client_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(), model_name: "gemini-files".to_string(),
model_name: "gemini-files".to_string(), mapped_model: candidate.selected_provider_model_name.clone(),
mapped_model: candidate.selected_provider_model_name.clone(), prompt_cache_key: None,
prompt_cache_key: None, provider_request_headers,
provider_request_headers, provider_request_body,
provider_request_body, provider_request_body_base64,
provider_request_body_base64, content_type: effective_headers
content_type: parts .get(http::header::CONTENT_TYPE)
.headers .and_then(|value| value.to_str().ok())
.get(http::header::CONTENT_TYPE) .map(str::trim)
.and_then(|value| value.to_str().ok()) .filter(|value| !value.is_empty())
.map(str::trim) .map(ToOwned::to_owned),
.filter(|value| !value.is_empty()) proxy,
.map(ToOwned::to_owned), transport_profile,
proxy, timeouts: resolve_transport_execution_timeouts(&transport),
transport_profile, upstream_is_stream: spec_metadata.require_streaming,
timeouts: resolve_transport_execution_timeouts(&transport), report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
upstream_is_stream: spec_metadata.require_streaming, report_context: Some(report_context),
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned), auth_context: input.auth_context.clone(),
report_context: Some(report_context), });
auth_context: input.auth_context.clone(), apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
}, Ok(Some(decision))
))
} }

View File

@@ -45,6 +45,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
let spec_metadata = local_gemini_files_spec_metadata(spec); let spec_metadata = local_gemini_files_spec_metadata(spec);
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let effective_headers = input.effective_headers(&parts.headers);
if let Some(skip_reason) = if let Some(skip_reason) =
gemini_files_transport_unsupported_reason(transport, GEMINI_FILES_CANDIDATE_API_FORMAT) gemini_files_transport_unsupported_reason(transport, GEMINI_FILES_CANDIDATE_API_FORMAT)
@@ -103,7 +104,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
body_is_empty, body_is_empty,
spec_metadata.decision_kind == GEMINI_FILES_UPLOAD_PLAN_KIND, spec_metadata.decision_kind == GEMINI_FILES_UPLOAD_PLAN_KIND,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
) { ) {
Ok(parts) => parts, Ok(parts) => parts,
Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload) => { Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload) => {
@@ -145,7 +146,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
}; };
let Some(provider_request_headers) = build_gemini_files_headers(GeminiFilesHeadersInput { let Some(provider_request_headers) = build_gemini_files_headers(GeminiFilesHeadersInput {
headers: &parts.headers, headers: effective_headers,
auth_header: &auth_header, auth_header: &auth_header,
auth_value: &auth_value, auth_value: &auth_value,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),

View File

@@ -14,7 +14,8 @@ use crate::ai_serving::planner::candidate_metadata::{
build_local_execution_candidate_metadata_for_candidate, LocalExecutionCandidateMetadataParts, build_local_execution_candidate_metadata_for_candidate, LocalExecutionCandidateMetadataParts,
}; };
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
build_local_authenticated_decision_input, resolve_local_authenticated_decision_input, attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind, build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
@@ -29,11 +30,12 @@ use crate::{AppState, GatewayError};
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalGeminiFilesCandidateAttempt; pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalGeminiFilesCandidateAttempt;
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalGeminiFilesCandidateAttemptSource; pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalGeminiFilesCandidateAttemptSource;
pub(super) use crate::ai_serving::planner::decision_input::LocalAuthenticatedDecisionInput as LocalGeminiFilesDecisionInput; pub(super) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalGeminiFilesDecisionInput;
pub(super) const GEMINI_FILES_CANDIDATE_API_FORMAT: &str = "gemini:files"; pub(super) const GEMINI_FILES_CANDIDATE_API_FORMAT: &str = "gemini:files";
pub(super) const GEMINI_FILES_CLIENT_API_FORMAT: &str = "gemini:files"; pub(super) const GEMINI_FILES_CLIENT_API_FORMAT: &str = "gemini:files";
pub(super) const GEMINI_FILES_REQUIRED_CAPABILITY: &str = "gemini_files"; pub(super) const GEMINI_FILES_REQUIRED_CAPABILITY: &str = "gemini_files";
pub(super) const GEMINI_FILES_ROUTING_MODEL: &str = "gemini-files";
pub(super) async fn resolve_local_gemini_files_decision_input( pub(super) async fn resolve_local_gemini_files_decision_input(
state: &AppState, state: &AppState,
@@ -41,9 +43,9 @@ pub(super) async fn resolve_local_gemini_files_decision_input(
body_json: Option<&serde_json::Value>, body_json: Option<&serde_json::Value>,
trace_id: &str, trace_id: &str,
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
) -> Option<LocalGeminiFilesDecisionInput> { ) -> Result<Option<LocalGeminiFilesDecisionInput>, GatewayError> {
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
return None; return Ok(None);
}; };
let explicit_required_capabilities = json!({ "gemini_files": true }); let explicit_required_capabilities = json!({ "gemini_files": true });
@@ -56,20 +58,33 @@ pub(super) async fn resolve_local_gemini_files_decision_input(
.await .await
{ {
Ok(Some(resolved_input)) => resolved_input, Ok(Some(resolved_input)) => resolved_input,
Ok(None) => return None, Ok(None) => return Ok(None),
Err(err) => { Err(err) => {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
error = ?err, error = ?err,
"gateway local gemini files decision auth snapshot read failed" "gateway local gemini files decision auth snapshot read failed"
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_authenticated_decision_input(resolved_input); let routing_body_json = body_json.cloned().unwrap_or(serde_json::Value::Null);
let mut input = build_local_requested_model_decision_input(
resolved_input,
GEMINI_FILES_ROUTING_MODEL.to_string(),
);
input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, body_json); input.client_session_affinity = client_session_affinity_from_parts(parts, body_json);
Some(input) attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
&routing_body_json,
GEMINI_FILES_CLIENT_API_FORMAT,
)
.await?;
Ok(Some(input))
} }
pub(super) async fn materialize_local_gemini_files_candidate_attempts( pub(super) async fn materialize_local_gemini_files_candidate_attempts(
@@ -101,8 +116,9 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
None, None,
None, input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
candidates, candidates,
Vec::new(), Vec::new(),
@@ -173,8 +189,9 @@ pub(super) async fn build_local_gemini_files_candidate_attempt_source<'a>(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
None, None,
None, input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
candidates, candidates,
Vec::new(), Vec::new(),

View File

@@ -32,7 +32,7 @@ pub(super) use crate::ai_serving::LocalOpenAiImageSpec;
pub(crate) struct LocalOpenAiImageSyncAttemptSource<'a> { pub(crate) struct LocalOpenAiImageSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
body_base64: Option<&'a str>, body_base64: Option<&'a str>,
trace_id: &'a str, trace_id: &'a str,
input: LocalOpenAiImageDecisionInput, input: LocalOpenAiImageDecisionInput,
@@ -43,7 +43,7 @@ pub(crate) struct LocalOpenAiImageSyncAttemptSource<'a> {
pub(crate) struct LocalOpenAiImageStreamAttemptSource<'a> { pub(crate) struct LocalOpenAiImageStreamAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
body_base64: Option<&'a str>, body_base64: Option<&'a str>,
trace_id: &'a str, trace_id: &'a str,
input: LocalOpenAiImageDecisionInput, input: LocalOpenAiImageDecisionInput,
@@ -152,16 +152,17 @@ pub(crate) async fn build_local_image_sync_attempt_source_for_kind<'a>(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source( let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
state, state,
trace_id, trace_id,
&input, &input,
body_json, &effective_body_json,
spec_metadata.api_format, spec_metadata.api_format,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
@@ -178,7 +179,7 @@ pub(crate) async fn build_local_image_sync_attempt_source_for_kind<'a>(
LocalOpenAiImageSyncAttemptSource { LocalOpenAiImageSyncAttemptSource {
state, state,
parts, parts,
body_json, body_json: effective_body_json,
body_base64, body_base64,
trace_id, trace_id,
input, input,
@@ -211,16 +212,17 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source( let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
state, state,
trace_id, trace_id,
&input, &input,
body_json, &effective_body_json,
spec_metadata.api_format, spec_metadata.api_format,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
@@ -237,7 +239,7 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
LocalOpenAiImageStreamAttemptSource { LocalOpenAiImageStreamAttemptSource {
state, state,
parts, parts,
body_json, body_json: effective_body_json,
body_base64, body_base64,
trace_id, trace_id,
input, input,
@@ -303,21 +305,21 @@ impl LocalOpenAiImageSyncAttemptSource<'_> {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate( let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
self.state, self.state,
self.parts, self.parts,
self.body_json, &self.body_json,
self.body_base64, self.body_base64,
self.trace_id, self.trace_id,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default(); let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default();
let built = if provider_api_format == "gemini:generate_content" { let built = if provider_api_format == "gemini:generate_content" {
build_gemini_sync_plan_from_decision(self.parts, self.body_json, payload) build_gemini_sync_plan_from_decision(self.parts, &self.body_json, payload)
} else { } else {
build_passthrough_sync_plan_from_decision(self.parts, payload) build_passthrough_sync_plan_from_decision(self.parts, payload)
}; };
@@ -345,23 +347,23 @@ impl LocalOpenAiImageStreamAttemptSource<'_> {
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate( let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
self.state, self.state,
self.parts, self.parts,
self.body_json, &self.body_json,
self.body_base64, self.body_base64,
self.trace_id, self.trace_id,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default(); let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default();
let built = if provider_api_format == "gemini:generate_content" { let built = if provider_api_format == "gemini:generate_content" {
build_gemini_stream_plan_from_decision(self.parts, self.body_json, payload) build_gemini_stream_plan_from_decision(self.parts, &self.body_json, payload)
} else { } else {
build_standard_stream_plan_from_decision(self.parts, self.body_json, payload, false) build_standard_stream_plan_from_decision(self.parts, &self.body_json, payload, false)
}; };
match built { match built {
Ok(value) => Ok(value), Ok(value) => Ok(value),
@@ -400,10 +402,11 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source( let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state, state,
@@ -429,7 +432,7 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
attempt, attempt,
spec, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -460,10 +463,11 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source( let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state, state,
@@ -489,7 +493,7 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
attempt, attempt,
spec, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -516,10 +520,11 @@ async fn build_local_sync_plan_and_reports(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source( let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state, state,
@@ -546,7 +551,7 @@ async fn build_local_sync_plan_and_reports(
attempt, attempt,
spec, spec,
) )
.await .await?
else { else {
continue; continue;
}; };
@@ -592,10 +597,11 @@ async fn build_local_stream_plan_and_reports(
trace_id, trace_id,
decision, decision,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source( let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
state, state,
@@ -622,7 +628,7 @@ async fn build_local_stream_plan_and_reports(
attempt, attempt,
spec, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -1,4 +1,5 @@
use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::report_context::{ use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, LocalExecutionReportContextParts, build_local_execution_report_context, LocalExecutionReportContextParts,
}; };
@@ -10,7 +11,9 @@ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_execution_timeouts, resolve_transport_profile,
}; };
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState}; use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{append_execution_contract_fields_to_value, AiExecutionDecision, AppState}; use crate::{
append_execution_contract_fields_to_value, AiExecutionDecision, AppState, GatewayError,
};
use super::request::resolve_local_openai_image_candidate_payload_parts; use super::request::resolve_local_openai_image_candidate_payload_parts;
use super::support::{LocalOpenAiImageCandidateAttempt, LocalOpenAiImageDecisionInput}; use super::support::{LocalOpenAiImageCandidateAttempt, LocalOpenAiImageDecisionInput};
@@ -25,11 +28,11 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
input: &LocalOpenAiImageDecisionInput, input: &LocalOpenAiImageDecisionInput,
attempt: LocalOpenAiImageCandidateAttempt, attempt: LocalOpenAiImageCandidateAttempt,
spec: LocalOpenAiImageSpec, spec: LocalOpenAiImageSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_openai_image_spec_metadata(spec); let spec_metadata = local_openai_image_spec_metadata(spec);
let planner_state = PlannerAppState::new(state); let planner_state = PlannerAppState::new(state);
let attempt_identity = attempt.attempt_identity(); let attempt_identity = attempt.attempt_identity();
let resolved = resolve_local_openai_image_candidate_payload_parts( let Some(resolved) = resolve_local_openai_image_candidate_payload_parts(
state, state,
parts, parts,
body_json, body_json,
@@ -39,7 +42,10 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
&attempt, &attempt,
spec, spec,
) )
.await?; .await
else {
return Ok(None);
};
let LocalOpenAiImageCandidateAttempt { let LocalOpenAiImageCandidateAttempt {
eligible, eligible,
candidate_id, candidate_id,
@@ -57,7 +63,10 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
.app() .app()
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport) .resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
.await; .await;
let transport_profile = resolve_transport_profile(&transport); let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&transport));
let mut extra_fields = serde_json::Map::new(); let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) { if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value); extra_fields.insert("proxy".to_string(), proxy_value);
@@ -88,6 +97,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
.get("stream") .get("stream")
.and_then(serde_json::Value::as_bool) .and_then(serde_json::Value::as_bool)
.unwrap_or(spec_metadata.require_streaming); .unwrap_or(spec_metadata.require_streaming);
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_execution_contract_fields_to_value( let report_context = append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts { build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context, auth_context: &input.auth_context,
@@ -114,7 +124,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
body_rules: transport.endpoint.body_rules.as_ref(), body_rules: transport.endpoint.body_rules.as_ref(),
provider_request_method: Some(serde_json::Value::String(parts.method.to_string())), provider_request_method: Some(serde_json::Value::String(parts.method.to_string())),
provider_request_headers: Some(&resolved.provider_request_headers), provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -134,39 +144,39 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
provider_api_format.as_str(), provider_api_format.as_str(),
); );
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: spec_metadata.require_streaming,
decision_is_stream: spec_metadata.require_streaming, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.clone(),
candidate_id: candidate_id.clone(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url: resolved.upstream_url,
upstream_url: resolved.upstream_url, provider_request_method: Some(parts.method.to_string()),
provider_request_method: Some(parts.method.to_string()), auth_header: Some(resolved.auth_header),
auth_header: Some(resolved.auth_header), auth_value: Some(resolved.auth_value),
auth_value: Some(resolved.auth_value), provider_api_format,
provider_api_format, client_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(), model_name: resolved.requested_model,
model_name: resolved.requested_model, mapped_model: resolved.mapped_model,
mapped_model: resolved.mapped_model, prompt_cache_key: None,
prompt_cache_key: None, provider_request_headers: resolved.provider_request_headers,
provider_request_headers: resolved.provider_request_headers, provider_request_body: Some(resolved.provider_request_body),
provider_request_body: Some(resolved.provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: Some("application/json".to_string()),
content_type: Some("application/json".to_string()), proxy,
proxy, transport_profile,
transport_profile, timeouts: resolve_transport_execution_timeouts(&transport),
timeouts: resolve_transport_execution_timeouts(&transport), upstream_is_stream,
upstream_is_stream, report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
)) Ok(Some(decision))
} }

View File

@@ -1,17 +1,19 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value; use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{ use crate::ai_serving::planner::candidate_preparation::{
prepare_header_authenticated_candidate, OauthPreparationContext, prepare_header_authenticated_candidate, OauthPreparationContext,
}; };
use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata; use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata;
use crate::ai_serving::pure::normalize_openai_image_request_with_options;
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_openai_image_headers, build_openai_image_upstream_url, build_grok_browser_headers, build_grok_upstream_url, build_openai_image_headers,
build_standard_provider_request_headers, openai_image_transport_unsupported_reason, build_openai_image_upstream_url, build_standard_provider_request_headers,
resolve_openai_image_auth, ProviderOpenAiImageHeadersInput, openai_image_transport_unsupported_reason, resolve_openai_image_auth, GrokHeaderInput,
StandardProviderRequestHeadersInput, ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers, apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
@@ -21,6 +23,7 @@ use crate::ai_serving::{
normalize_openai_image_request, request_conversion_direct_auth, CandidateFailureDiagnostic, normalize_openai_image_request, request_conversion_direct_auth, CandidateFailureDiagnostic,
GatewayProviderTransportSnapshot, PlannerAppState, RequestConversionKind, GatewayProviderTransportSnapshot, PlannerAppState, RequestConversionKind,
}; };
use crate::image_capabilities::openai_image_normalize_options_for_provider;
use crate::AppState; use crate::AppState;
use super::support::{ use super::support::{
@@ -43,6 +46,7 @@ pub(super) struct LocalOpenAiImageCandidatePayloadParts {
pub(super) provider_request_body: Value, pub(super) provider_request_body: Value,
pub(super) upstream_url: String, pub(super) upstream_url: String,
pub(super) input_summary: Value, pub(super) input_summary: Value,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
} }
pub(super) async fn resolve_local_openai_image_candidate_payload_parts( pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
@@ -59,6 +63,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let provider_api_format = attempt.eligible.provider_api_format.as_str(); let provider_api_format = attempt.eligible.provider_api_format.as_str();
let effective_headers = input.effective_headers(&parts.headers);
if provider_api_format == "gemini:generate_content" { if provider_api_format == "gemini:generate_content" {
return resolve_local_openai_image_to_gemini_candidate_payload_parts( return resolve_local_openai_image_to_gemini_candidate_payload_parts(
@@ -120,8 +125,13 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
let auth_header = prepared_candidate.auth_header; let auth_header = prepared_candidate.auth_header;
let auth_value = prepared_candidate.auth_value; let auth_value = prepared_candidate.auth_value;
let Some(normalized_request) = normalize_openai_image_request(parts, body_json, body_base64) let normalized_request = normalize_openai_image_request_with_options(
else { parts,
body_json,
body_base64,
openai_image_normalize_options_for_provider(&transport.provider.provider_type),
);
let Some(normalized_request) = normalized_request else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic( mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state, state,
input, input,
@@ -145,8 +155,16 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.provider_type .provider_type
.trim() .trim()
.eq_ignore_ascii_case("chatgpt_web"); .eq_ignore_ascii_case("chatgpt_web");
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let upstream_url = if is_chatgpt_web { let upstream_url = if is_chatgpt_web {
chatgpt_web_image_internal_url(&transport.endpoint.base_url) chatgpt_web_image_internal_url(&transport.endpoint.base_url)
} else if is_grok {
build_grok_upstream_url(transport, GROK_CHAT_PATH)
} else { } else {
build_openai_image_upstream_url(transport, parts.uri.query()) build_openai_image_upstream_url(transport, parts.uri.query())
}; };
@@ -168,16 +186,27 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
); );
} }
let Some(mut provider_request_headers) = let Some(mut provider_request_headers) = (if is_grok {
build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "*/*",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
})
} else {
build_openai_image_headers(ProviderOpenAiImageHeadersInput { build_openai_image_headers(ProviderOpenAiImageHeadersInput {
headers: &parts.headers, headers: effective_headers,
auth_header: &auth_header, auth_header: &auth_header,
auth_value: &auth_value, auth_value: &auth_value,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: body_json, original_request_body: body_json,
}) })
else { }) else {
mark_skipped_local_openai_image_candidate_with_failure_diagnostic( mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
state, state,
input, input,
@@ -197,11 +226,12 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
}; };
if is_chatgpt_web { if is_chatgpt_web {
provider_request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string()); provider_request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string());
} else if is_grok {
} else { } else {
apply_codex_openai_responses_special_headers( apply_codex_openai_responses_special_headers(
&mut provider_request_headers, &mut provider_request_headers,
&provider_request_body, &provider_request_body,
&parts.headers, effective_headers,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
spec_metadata.api_format, spec_metadata.api_format,
Some(trace_id), Some(trace_id),
@@ -222,7 +252,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
.unwrap_or_default() .unwrap_or_default()
.to_string(); .to_string();
let input_summary = if is_chatgpt_web { let input_summary = if is_chatgpt_web || is_grok {
provider_request_body.clone() provider_request_body.clone()
} else { } else {
normalized_request.summary_json normalized_request.summary_json
@@ -239,6 +269,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
provider_request_body, provider_request_body,
upstream_url, upstream_url,
input_summary, input_summary,
transport_profile,
}) })
} }
@@ -256,6 +287,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let provider_api_format = "gemini:generate_content"; let provider_api_format = "gemini:generate_content";
let effective_headers = input.effective_headers(&parts.headers);
let prepared_candidate = match prepare_header_authenticated_candidate( let prepared_candidate = match prepare_header_authenticated_candidate(
PlannerAppState::new(state), PlannerAppState::new(state),
@@ -332,7 +364,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
converted.body_json, converted.body_json,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
body_json, body_json,
&parts.headers, effective_headers,
) { ) {
Some(body) => body, Some(body) => body,
None => { None => {
@@ -385,7 +417,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
transport, transport,
provider_api_format, provider_api_format,
same_format: false, same_format: false,
headers: &parts.headers, headers: effective_headers,
auth_header: &prepared_candidate.auth_header, auth_header: &prepared_candidate.auth_header,
auth_value: &prepared_candidate.auth_value, auth_value: &prepared_candidate.auth_value,
extra_headers: &BTreeMap::new(), extra_headers: &BTreeMap::new(),
@@ -424,6 +456,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
provider_request_body: converted.body_json, provider_request_body: converted.body_json,
upstream_url, upstream_url,
input_summary: converted.summary_json, input_summary: converted.summary_json,
transport_profile: None,
}) })
} }

View File

@@ -13,6 +13,7 @@ use crate::ai_serving::planner::candidate_metadata::{
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate; use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::planner::candidate_source::auth_snapshot_allows_cross_format_candidate; use crate::ai_serving::planner::candidate_source::auth_snapshot_allows_cross_format_candidate;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
@@ -42,12 +43,16 @@ pub(super) async fn resolve_local_openai_image_decision_input(
body_base64: Option<&str>, body_base64: Option<&str>,
trace_id: &str, trace_id: &str,
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
) -> Option<LocalOpenAiImageDecisionInput> { ) -> Result<Option<LocalOpenAiImageDecisionInput>, GatewayError> {
let Some(auth_context) = resolve_local_openai_image_auth_context(decision) else { let Some(auth_context) = resolve_local_openai_image_auth_context(decision) else {
return None; return Ok(None);
}; };
let requested_model = resolve_requested_image_model_for_request(parts, body_json, body_base64)?; let Some(requested_model) =
resolve_requested_image_model_for_request(parts, body_json, body_base64)
else {
return Ok(None);
};
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
state, state,
@@ -58,21 +63,37 @@ pub(super) async fn resolve_local_openai_image_decision_input(
.await .await
{ {
Ok(Some(resolved_input)) => resolved_input, Ok(Some(resolved_input)) => resolved_input,
Ok(None) => return None, Ok(None) => return Ok(None),
Err(err) => { Err(err) => {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
error = ?err, error = ?err,
"gateway local openai image decision auth snapshot read failed" "gateway local openai image decision auth snapshot read failed"
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
"openai:image",
)
.await
{
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai image decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }
fn resolve_local_openai_image_auth_context( fn resolve_local_openai_image_auth_context(
@@ -229,6 +250,7 @@ pub(super) async fn build_local_openai_image_candidate_attempt_source<'a>(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -305,6 +327,7 @@ async fn materialize_local_openai_image_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,

View File

@@ -27,7 +27,7 @@ use self::support::{
pub(crate) struct LocalVideoCreateSyncAttemptSource<'a> { pub(crate) struct LocalVideoCreateSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
trace_id: &'a str, trace_id: &'a str,
input: LocalVideoCreateDecisionInput, input: LocalVideoCreateDecisionInput,
spec: LocalVideoCreateSpec, spec: LocalVideoCreateSpec,
@@ -65,16 +65,17 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
let Some(input) = resolve_local_video_create_decision_input( let Some(input) = resolve_local_video_create_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let Some((candidates, candidate_count)) = build_local_video_create_candidate_attempt_source( let Some((candidates, candidate_count)) = build_local_video_create_candidate_attempt_source(
state, state,
trace_id, trace_id,
&input, &input,
body_json, &effective_body_json,
spec_metadata.api_format, spec_metadata.api_format,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
@@ -91,7 +92,7 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
LocalVideoCreateSyncAttemptSource { LocalVideoCreateSyncAttemptSource {
state, state,
parts, parts,
body_json, body_json: effective_body_json,
trace_id, trace_id,
input, input,
spec, spec,
@@ -133,13 +134,13 @@ impl LocalVideoCreateSyncAttemptSource<'_> {
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate( let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
self.state, self.state,
self.parts, self.parts,
self.body_json, &self.body_json,
self.trace_id, self.trace_id,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -175,10 +176,11 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
let Some(input) = resolve_local_video_create_decision_input( let Some(input) = resolve_local_video_create_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source( let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
state, state,
@@ -197,7 +199,7 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate( if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec, state, parts, body_json, trace_id, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -218,10 +220,11 @@ async fn build_local_sync_plan_and_reports(
let Some(input) = resolve_local_video_create_decision_input( let Some(input) = resolve_local_video_create_decision_input(
state, parts, trace_id, decision, body_json, spec, state, parts, trace_id, decision, body_json, spec,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
let body_json = input.effective_body_json(body_json);
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source( let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
state, state,
@@ -241,7 +244,7 @@ async fn build_local_sync_plan_and_reports(
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate( let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
state, parts, body_json, trace_id, &input, attempt, spec, state, parts, body_json, trace_id, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -1,4 +1,5 @@
use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::report_context::{ use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, LocalExecutionReportContextParts, build_local_execution_report_context, LocalExecutionReportContextParts,
}; };
@@ -10,7 +11,7 @@ use crate::ai_serving::transport::{
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_execution_timeouts, resolve_transport_profile,
}; };
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState}; use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
use crate::{AiExecutionDecision, AppState}; use crate::{AiExecutionDecision, AppState, GatewayError};
use super::request::resolve_local_video_create_candidate_payload_parts; use super::request::resolve_local_video_create_candidate_payload_parts;
use super::support::{LocalVideoCreateCandidateAttempt, LocalVideoCreateDecisionInput}; use super::support::{LocalVideoCreateCandidateAttempt, LocalVideoCreateDecisionInput};
@@ -24,14 +25,17 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
input: &LocalVideoCreateDecisionInput, input: &LocalVideoCreateDecisionInput,
attempt: LocalVideoCreateCandidateAttempt, attempt: LocalVideoCreateCandidateAttempt,
spec: LocalVideoCreateSpec, spec: LocalVideoCreateSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_video_create_spec_metadata(spec); let spec_metadata = local_video_create_spec_metadata(spec);
let planner_state = PlannerAppState::new(state); let planner_state = PlannerAppState::new(state);
let attempt_identity = attempt.attempt_identity(); let attempt_identity = attempt.attempt_identity();
let resolved = resolve_local_video_create_candidate_payload_parts( let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
state, parts, body_json, trace_id, input, &attempt, spec, state, parts, body_json, trace_id, input, &attempt, spec,
) )
.await?; .await
else {
return Ok(None);
};
let LocalVideoCreateCandidateAttempt { let LocalVideoCreateCandidateAttempt {
eligible, eligible,
candidate_id, candidate_id,
@@ -50,6 +54,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) { if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value); extra_fields.insert("proxy".to_string(), proxy_value);
} }
let effective_headers = input.effective_headers(&parts.headers);
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts { let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context, auth_context: &input.auth_context,
request_id: trace_id, request_id: trace_id,
@@ -75,7 +80,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
body_rules: transport.endpoint.body_rules.as_ref(), body_rules: transport.endpoint.body_rules.as_ref(),
provider_request_method: None, provider_request_method: None,
provider_request_headers: None, provider_request_headers: None,
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -99,45 +104,45 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
upstream_url, upstream_url,
} = resolved; } = resolved;
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: false,
decision_is_stream: false, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.clone(),
candidate_id: candidate_id.clone(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: Some(parts.method.to_string()),
provider_request_method: Some(parts.method.to_string()), auth_header: Some(auth_header),
auth_header: Some(auth_header), auth_value: Some(auth_value),
auth_value: Some(auth_value), provider_api_format: spec_metadata.api_format.to_string(),
provider_api_format: spec_metadata.api_format.to_string(), client_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(), model_name: input.requested_model.clone(),
model_name: input.requested_model.clone(), mapped_model,
mapped_model, prompt_cache_key: None,
prompt_cache_key: None, provider_request_headers,
provider_request_headers, provider_request_body: Some(provider_request_body),
provider_request_body: Some(provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: parts
content_type: parts .headers
.headers .get(http::header::CONTENT_TYPE)
.get(http::header::CONTENT_TYPE) .and_then(|value| value.to_str().ok())
.and_then(|value| value.to_str().ok()) .map(str::trim)
.map(str::trim) .filter(|value| !value.is_empty())
.filter(|value| !value.is_empty()) .map(ToOwned::to_owned),
.map(ToOwned::to_owned), proxy,
proxy, transport_profile,
transport_profile, timeouts: resolve_transport_execution_timeouts(&transport),
timeouts: resolve_transport_execution_timeouts(&transport), upstream_is_stream: false,
upstream_is_stream: false, report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
)) Ok(Some(decision))
} }

View File

@@ -41,6 +41,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
let spec_metadata = local_video_create_spec_metadata(spec); let spec_metadata = local_video_create_spec_metadata(spec);
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let effective_headers = input.effective_headers(&parts.headers);
let provider_family = provider_video_create_family(spec.family); let provider_family = provider_video_create_family(spec.family);
let transport_unsupported_reason = video_create_transport_unsupported_reason( let transport_unsupported_reason = video_create_transport_unsupported_reason(
@@ -124,7 +125,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
provider_family, provider_family,
&mapped_model, &mapped_model,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
) else { ) else {
mark_skipped_local_video_candidate_with_failure_diagnostic( mark_skipped_local_video_candidate_with_failure_diagnostic(
state, state,
@@ -146,7 +147,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
let Some(provider_request_headers) = let Some(provider_request_headers) =
build_video_create_headers(ProviderVideoCreateHeadersInput { build_video_create_headers(ProviderVideoCreateHeadersInput {
headers: &parts.headers, headers: effective_headers,
auth_header: &auth_header, auth_header: &auth_header,
auth_value: &auth_value, auth_value: &auth_value,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),

View File

@@ -15,6 +15,7 @@ use crate::ai_serving::planner::candidate_metadata::{
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate; use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
use crate::ai_serving::planner::common::extract_requested_model_from_request; use crate::ai_serving::planner::common::extract_requested_model_from_request;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
@@ -41,19 +42,21 @@ pub(super) async fn resolve_local_video_create_decision_input(
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
body_json: &serde_json::Value, body_json: &serde_json::Value,
spec: LocalVideoCreateSpec, spec: LocalVideoCreateSpec,
) -> Option<LocalVideoCreateDecisionInput> { ) -> Result<Option<LocalVideoCreateDecisionInput>, GatewayError> {
let spec_metadata = local_video_create_spec_metadata(spec); let spec_metadata = local_video_create_spec_metadata(spec);
let Some(auth_context) = resolve_local_video_create_auth_context(decision, spec.family) else { let Some(auth_context) = resolve_local_video_create_auth_context(decision, spec.family) else {
return None; return Ok(None);
}; };
let requested_model = extract_requested_model_from_request( let Some(requested_model) = extract_requested_model_from_request(
parts, parts,
body_json, body_json,
spec_metadata spec_metadata
.requested_model_family .requested_model_family
.expect("video specs should declare requested-model family"), .expect("video specs should declare requested-model family"),
)?; ) else {
return Ok(None);
};
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
state, state,
@@ -64,7 +67,7 @@ pub(super) async fn resolve_local_video_create_decision_input(
.await .await
{ {
Ok(Some(resolved_input)) => resolved_input, Ok(Some(resolved_input)) => resolved_input,
Ok(None) => return None, Ok(None) => return Ok(None),
Err(err) => { Err(err) => {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
@@ -72,14 +75,31 @@ pub(super) async fn resolve_local_video_create_decision_input(
error = ?err, error = ?err,
"gateway local video decision auth snapshot read failed" "gateway local video decision auth snapshot read failed"
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
spec_metadata.api_format,
)
.await
{
warn!(
trace_id = %trace_id,
decision_kind = spec_metadata.decision_kind,
error = ?err,
"gateway local video decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }
fn resolve_local_video_create_auth_context( fn resolve_local_video_create_auth_context(
@@ -196,6 +216,7 @@ pub(super) async fn build_local_video_create_candidate_attempt_source<'a>(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -261,6 +282,7 @@ async fn materialize_local_video_create_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,

View File

@@ -29,7 +29,7 @@ pub(crate) struct LocalStandardSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalStandardDecisionInput, input: LocalStandardDecisionInput,
spec: LocalStandardSpec, spec: LocalStandardSpec,
requested_model_family: RequestedModelFamily, requested_model_family: RequestedModelFamily,
@@ -40,7 +40,7 @@ pub(crate) struct LocalStandardStreamAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalStandardDecisionInput, input: LocalStandardDecisionInput,
spec: LocalStandardSpec, spec: LocalStandardSpec,
requested_model_family: RequestedModelFamily, requested_model_family: RequestedModelFamily,
@@ -61,7 +61,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
.expect("standard spec metadata should include requested-model family"); .expect("standard spec metadata should include requested-model family");
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -82,9 +82,15 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let (candidates, candidate_count) = let effective_body_json = input.effective_body_json(body_json).clone();
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) let (candidates, candidate_count) = build_local_standard_candidate_attempt_source(
.await?; state,
trace_id,
&input,
&effective_body_json,
spec,
)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count); apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
if candidate_count == 0 { if candidate_count == 0 {
return Ok(None); return Ok(None);
@@ -95,7 +101,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
requested_model_family, requested_model_family,
@@ -119,7 +125,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
.expect("standard spec metadata should include requested-model family"); .expect("standard spec metadata should include requested-model family");
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -140,9 +146,15 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let (candidates, candidate_count) = let effective_body_json = input.effective_body_json(body_json).clone();
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) let (candidates, candidate_count) = build_local_standard_candidate_attempt_source(
.await?; state,
trace_id,
&input,
&effective_body_json,
spec,
)
.await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count); apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
if candidate_count == 0 { if candidate_count == 0 {
return Ok(None); return Ok(None);
@@ -153,7 +165,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
requested_model_family, requested_model_family,
@@ -228,19 +240,19 @@ impl LocalStandardSyncAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
match build_sync_plan_from_requested_model_family( match build_sync_plan_from_requested_model_family(
self.requested_model_family, self.requested_model_family,
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
) { ) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
@@ -265,19 +277,19 @@ impl LocalStandardStreamAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
match build_stream_plan_from_requested_model_family( match build_stream_plan_from_requested_model_family(
self.requested_model_family, self.requested_model_family,
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
) { ) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
@@ -309,7 +321,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -322,6 +334,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?; .await?;
@@ -331,7 +344,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate( if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -358,7 +371,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -371,6 +384,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?; .await?;
@@ -380,7 +394,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate( if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -405,7 +419,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
.expect("standard spec metadata should include requested-model family"); .expect("standard spec metadata should include requested-model family");
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -426,6 +440,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?; .await?;
@@ -438,7 +453,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate( let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };
@@ -479,7 +494,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
.expect("standard spec metadata should include requested-model family"); .expect("standard spec metadata should include requested-model family");
let Some(input) = let Some(input) =
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec) resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
.await .await?
else { else {
set_local_runtime_miss_diagnostic_reason( set_local_runtime_miss_diagnostic_reason(
state, state,
@@ -500,6 +515,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let body_json = input.effective_body_json(body_json);
let (mut source, candidate_count) = let (mut source, candidate_count) =
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec) build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
.await?; .await?;
@@ -512,7 +528,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate( let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -16,6 +16,7 @@ use crate::ai_serving::planner::candidate_source::{
}; };
use crate::ai_serving::planner::common::extract_requested_model_from_request; use crate::ai_serving::planner::common::extract_requested_model_from_request;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
@@ -39,19 +40,21 @@ pub(super) async fn resolve_local_standard_decision_input(
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
body_json: &serde_json::Value, body_json: &serde_json::Value,
spec: LocalStandardSpec, spec: LocalStandardSpec,
) -> Option<LocalStandardDecisionInput> { ) -> Result<Option<LocalStandardDecisionInput>, GatewayError> {
let spec_metadata = local_standard_spec_metadata(spec); let spec_metadata = local_standard_spec_metadata(spec);
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
return None; return Ok(None);
}; };
let requested_model = extract_requested_model_from_request( let Some(requested_model) = extract_requested_model_from_request(
parts, parts,
body_json, body_json,
spec_metadata spec_metadata
.requested_model_family .requested_model_family
.expect("standard specs should declare requested-model family"), .expect("standard specs should declare requested-model family"),
)?; ) else {
return Ok(None);
};
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
state, state,
@@ -62,7 +65,7 @@ pub(super) async fn resolve_local_standard_decision_input(
.await .await
{ {
Ok(Some(resolved_input)) => resolved_input, Ok(Some(resolved_input)) => resolved_input,
Ok(None) => return None, Ok(None) => return Ok(None),
Err(err) => { Err(err) => {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
@@ -70,14 +73,31 @@ pub(super) async fn resolve_local_standard_decision_input(
error = ?err, error = ?err,
"gateway local standard decision auth snapshot read failed" "gateway local standard decision auth snapshot read failed"
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
spec_metadata.api_format,
)
.await
{
warn!(
trace_id = %trace_id,
api_format = spec_metadata.api_format,
error = ?err,
"gateway local standard decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }
pub(super) async fn materialize_local_standard_candidate_attempts( pub(super) async fn materialize_local_standard_candidate_attempts(
@@ -104,6 +124,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
spec_metadata.require_streaming, spec_metadata.require_streaming,
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
&input.auth_snapshot, &input.auth_snapshot,
input.routing_policy.as_ref(),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
false, false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
@@ -128,6 +149,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -228,6 +250,7 @@ pub(super) async fn build_local_standard_candidate_attempt_source<'a>(
&input.auth_snapshot, &input.auth_snapshot,
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -320,6 +343,7 @@ async fn maybe_append_gemini_image_openai_image_preselection(
spec_metadata.require_streaming, spec_metadata.require_streaming,
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
&input.auth_snapshot, &input.auth_snapshot,
input.routing_policy.as_ref(),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
false, false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,

View File

@@ -3,6 +3,7 @@ use crate::ai_serving::planner::candidate_materialization::{
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data, mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
mark_skipped_local_execution_candidate_with_failure_diagnostic, mark_skipped_local_execution_candidate_with_failure_diagnostic,
}; };
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind, build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
}; };
@@ -24,7 +25,7 @@ use crate::ai_serving::{
}; };
use crate::{ use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value, append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, AiExecutionDecision, AppState, GatewayError,
}; };
use super::request::resolve_local_standard_candidate_payload_parts; use super::request::resolve_local_standard_candidate_payload_parts;
@@ -38,7 +39,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
input: &LocalStandardDecisionInput, input: &LocalStandardDecisionInput,
attempt: LocalStandardCandidateAttempt, attempt: LocalStandardCandidateAttempt,
spec: LocalStandardSpec, spec: LocalStandardSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_standard_spec_metadata(spec); let spec_metadata = local_standard_spec_metadata(spec);
if api_format_alias_matches( if api_format_alias_matches(
&attempt.eligible.provider_api_format, &attempt.eligible.provider_api_format,
@@ -70,10 +71,13 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
.. ..
} = &attempt; } = &attempt;
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let resolved = resolve_local_standard_candidate_payload_parts( let Some(resolved) = resolve_local_standard_candidate_payload_parts(
state, parts, trace_id, body_json, input, &attempt, spec, state, parts, trace_id, body_json, input, &attempt, spec,
) )
.await?; .await
else {
return Ok(None);
};
let proxy = state let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport) .resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await; .await;
@@ -93,6 +97,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
spec_metadata.api_format, spec_metadata.api_format,
resolved.provider_api_format.as_str(), resolved.provider_api_format.as_str(),
); );
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_local_failover_policy_to_value( let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value( append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts { build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -120,7 +125,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
body_rules: resolved.transport.endpoint.body_rules.as_ref(), body_rules: resolved.transport.endpoint.body_rules.as_ref(),
provider_request_method: Some(serde_json::Value::Null), provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers), provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -144,7 +149,10 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
), ),
&resolved.transport, &resolved.transport,
); );
let transport_profile = resolve_transport_profile(&resolved.transport); let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport); let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let super::request::LocalStandardCandidatePayloadParts { let super::request::LocalStandardCandidatePayloadParts {
auth_header, auth_header,
@@ -157,43 +165,44 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
upstream_is_stream, upstream_is_stream,
envelope_name: _, envelope_name: _,
transport, transport,
transport_profile: _,
} = resolved; } = resolved;
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: spec_metadata.require_streaming,
decision_is_stream: spec_metadata.require_streaming, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.to_string(),
candidate_id: candidate_id.to_string(), provider_name: candidate.provider_name.clone(),
provider_name: candidate.provider_name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: None,
provider_request_method: None, auth_header: Some(auth_header),
auth_header: Some(auth_header), auth_value: Some(auth_value),
auth_value: Some(auth_value), provider_api_format,
provider_api_format, client_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(), model_name: input.requested_model.clone(),
model_name: input.requested_model.clone(), mapped_model,
mapped_model, prompt_cache_key: None,
prompt_cache_key: None, provider_request_headers,
provider_request_headers, provider_request_body: Some(provider_request_body),
provider_request_body: Some(provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: Some("application/json".to_string()),
content_type: Some("application/json".to_string()), proxy,
proxy, transport_profile,
transport_profile, timeouts,
timeouts, upstream_is_stream,
upstream_is_stream, report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
)) Ok(Some(decision))
} }
pub(super) async fn mark_skipped_local_standard_candidate( pub(super) async fn mark_skipped_local_standard_candidate(
@@ -347,6 +356,9 @@ mod tests {
required_capabilities: None, required_capabilities: None,
request_auth_channel: None, request_auth_channel: None,
client_session_affinity: None, client_session_affinity: None,
routing_policy: None,
routing_trace_seed: None,
routing_context: None,
} }
} }
@@ -512,6 +524,7 @@ mod tests {
claude_stream_spec(), claude_stream_spec(),
) )
.await .await
.expect("same-format candidate should not fail routing mutation")
.expect("same-format candidate should build a standard-family payload"); .expect("same-format candidate should build a standard-family payload");
assert_eq!(payload.endpoint_id.as_deref(), Some("endpoint-claude")); assert_eq!(payload.endpoint_id.as_deref(), Some("endpoint-claude"));
@@ -547,6 +560,7 @@ mod tests {
claude_stream_spec(), claude_stream_spec(),
) )
.await .await
.expect("cross-format candidate should not fail routing mutation")
.expect("cross-format candidate should still build after the same-format candidate"); .expect("cross-format candidate should still build after the same-format candidate");
assert_eq!( assert_eq!(

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value; use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{ use crate::ai_serving::planner::candidate_preparation::{
@@ -21,10 +22,11 @@ use crate::ai_serving::transport::kiro::{
KIRO_ENVELOPE_NAME, KIRO_ENVELOPE_NAME,
}; };
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_openai_image_headers, build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
build_openai_image_upstream_url, build_standard_provider_request_headers, build_openai_image_headers, build_openai_image_upstream_url,
openai_image_transport_unsupported_reason, resolve_openai_image_auth, build_standard_provider_request_headers, openai_image_transport_unsupported_reason,
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, resolve_grok_session_auth, resolve_openai_image_auth, GrokHeaderInput,
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
build_openai_image_request_body_from_gemini_image_request, gemini_request_is_image_generation, build_openai_image_request_body_from_gemini_image_request, gemini_request_is_image_generation,
@@ -49,6 +51,14 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
pub(super) upstream_is_stream: bool, pub(super) upstream_is_stream: bool,
pub(super) envelope_name: Option<&'static str>, pub(super) envelope_name: Option<&'static str>,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>, pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
} }
pub(crate) async fn resolve_local_standard_candidate_payload_parts( pub(crate) async fn resolve_local_standard_candidate_payload_parts(
@@ -64,7 +74,14 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
let planner_state = crate::ai_serving::PlannerAppState::new(state); let planner_state = crate::ai_serving::PlannerAppState::new(state);
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let transport = &attempt.eligible.transport; let transport = &attempt.eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let provider_api_format = attempt.eligible.provider_api_format.as_str(); let provider_api_format = attempt.eligible.provider_api_format.as_str();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
if spec_metadata.api_format == "gemini:generate_content" if spec_metadata.api_format == "gemini:generate_content"
&& provider_api_format == "openai:image" && provider_api_format == "openai:image"
&& gemini_request_is_image_generation(body_json) && gemini_request_is_image_generation(body_json)
@@ -75,11 +92,116 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
.await; .await;
} }
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format); let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
if !crate::ai_serving::request_pair_allowed_for_transport( if is_grok && is_grok_text_provider_api_format(provider_api_format) {
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
candidate,
resolve_grok_session_auth(transport),
OauthPreparationContext {
trace_id,
api_format: provider_api_format,
operation: "standard_family_grok_text_request",
},
)
.await
{
Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_standard_candidate(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
skip_reason,
)
.await;
return None;
}
};
let mut provider_request_body = body_json.clone();
if let Some(object) = provider_request_body.as_object_mut() {
object.insert(
"model".to_string(),
serde_json::Value::String(prepared_candidate.mapped_model.clone()),
);
}
let upstream_is_stream = resolve_upstream_is_stream_for_provider(
transport.endpoint.config.as_ref(),
transport.provider.provider_type.as_str(),
provider_api_format,
spec_metadata.require_streaming,
false,
);
let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
enforce_provider_body_stream_policy(
&mut provider_request_body,
provider_api_format,
upstream_is_stream,
request_requires_body_stream_field(body_json, force_body_stream_field),
);
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
let Some(provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
}) else {
mark_skipped_local_standard_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
attempt.candidate_index,
&attempt.candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
provider_api_format,
"grok_standard_family_headers",
),
)
.await;
return None;
};
return Some(LocalStandardCandidatePayloadParts {
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
mapped_model: prepared_candidate.mapped_model,
provider_api_format: provider_api_format.to_string(),
provider_request_body,
provider_request_headers,
upstream_url,
upstream_is_stream,
envelope_name: None,
transport: Arc::clone(transport),
transport_profile,
});
}
let Some(conversion_kind) =
crate::ai_serving::request_conversion_kind(spec_metadata.api_format, provider_api_format)
else {
return None;
};
if crate::ai_serving::request_conversion_transport_unsupported_reason(
transport, transport,
spec_metadata.api_format, conversion_kind,
provider_api_format, )
) { .is_some()
{
return None; return None;
} }
@@ -212,7 +334,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
transport.endpoint.body_rules.as_ref() transport.endpoint.body_rules.as_ref()
}, },
Some(input.auth_context.api_key_id.as_str()), Some(input.auth_context.api_key_id.as_str()),
Some(&parts.headers), Some(effective_headers),
enable_model_directives, enable_model_directives,
) { ) {
Some(body) => body, Some(body) => body,
@@ -362,7 +484,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
transport, transport,
provider_api_format, provider_api_format,
same_format: false, same_format: false,
headers: &parts.headers, headers: effective_headers,
auth_header: &prepared_candidate.auth_header, auth_header: &prepared_candidate.auth_header,
auth_value: &prepared_candidate.auth_value, auth_value: &prepared_candidate.auth_value,
extra_headers: &BTreeMap::new(), extra_headers: &BTreeMap::new(),
@@ -393,7 +515,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
apply_codex_openai_responses_special_headers( apply_codex_openai_responses_special_headers(
&mut provider_request_headers, &mut provider_request_headers,
&provider_request_body, &provider_request_body,
&parts.headers, effective_headers,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
provider_api_format, provider_api_format,
Some(trace_id), Some(trace_id),
@@ -411,6 +533,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
upstream_is_stream, upstream_is_stream,
envelope_name: None, envelope_name: None,
transport: Arc::clone(transport), transport: Arc::clone(transport),
transport_profile: None,
}) })
} }
@@ -510,9 +633,10 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
let upstream_is_stream = true; let upstream_is_stream = true;
let upstream_url = build_openai_image_upstream_url(transport, None); let upstream_url = build_openai_image_upstream_url(transport, None);
let effective_headers = input.effective_headers(&parts.headers);
let Some(mut provider_request_headers) = let Some(mut provider_request_headers) =
build_openai_image_headers(ProviderOpenAiImageHeadersInput { build_openai_image_headers(ProviderOpenAiImageHeadersInput {
headers: &parts.headers, headers: effective_headers,
auth_header: &prepared_candidate.auth_header, auth_header: &prepared_candidate.auth_header,
auth_value: &prepared_candidate.auth_value, auth_value: &prepared_candidate.auth_value,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
@@ -540,7 +664,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
apply_codex_openai_responses_special_headers( apply_codex_openai_responses_special_headers(
&mut provider_request_headers, &mut provider_request_headers,
&converted.body_json, &converted.body_json,
&parts.headers, effective_headers,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
provider_api_format, provider_api_format,
Some(trace_id), Some(trace_id),
@@ -558,6 +682,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
upstream_is_stream, upstream_is_stream,
envelope_name: None, envelope_name: None,
transport: Arc::clone(transport), transport: Arc::clone(transport),
transport_profile: None,
}) })
} }
@@ -579,12 +704,13 @@ async fn build_kiro_cross_format_payload_parts(
kiro_auth: &KiroRequestAuth, kiro_auth: &KiroRequestAuth,
) -> Option<LocalStandardCandidatePayloadParts> { ) -> Option<LocalStandardCandidatePayloadParts> {
let candidate = &attempt.eligible.candidate; let candidate = &attempt.eligible.candidate;
let effective_headers = input.effective_headers(&parts.headers);
let provider_request_body = match build_kiro_provider_request_body( let provider_request_body = match build_kiro_provider_request_body(
&claude_request_body, &claude_request_body,
&mapped_model, &mapped_model,
&kiro_auth.auth_config, &kiro_auth.auth_config,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
) { ) {
Some(body) => body, Some(body) => body,
None => { None => {
@@ -635,7 +761,7 @@ async fn build_kiro_cross_format_payload_parts(
} }
}; };
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput { let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
headers: &parts.headers, headers: effective_headers,
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: original_body_json, original_request_body: original_body_json,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
@@ -676,5 +802,6 @@ async fn build_kiro_cross_format_payload_parts(
upstream_is_stream, upstream_is_stream,
envelope_name: Some(KIRO_ENVELOPE_NAME), envelope_name: Some(KIRO_ENVELOPE_NAME),
transport: Arc::clone(transport), transport: Arc::clone(transport),
transport_profile: None,
}) })
} }

View File

@@ -1,5 +1,6 @@
use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::common::OPENAI_CHAT_STREAM_PLAN_KIND; use crate::ai_serving::planner::common::OPENAI_CHAT_STREAM_PLAN_KIND;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::report_context::{ use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, insert_provider_stream_event_api_format, build_local_execution_report_context, insert_provider_stream_event_api_format,
LocalExecutionReportContextParts, LocalExecutionReportContextParts,
@@ -67,7 +68,10 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
let proxy = state let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport) .resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await; .await;
let transport_profile = resolve_transport_profile(&resolved.transport); let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport); let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let mut extra_fields = serde_json::Map::new(); let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) = if let Some(proxy_value) =
@@ -99,12 +103,14 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
envelope_name, envelope_name,
transport, transport,
request_redacted, request_redacted,
transport_profile: _,
} = resolved; } = resolved;
let original_request_body_json = if request_redacted { let original_request_body_json = if request_redacted {
Some(&provider_request_body) Some(&provider_request_body)
} else { } else {
Some(body_json) Some(body_json)
}; };
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_local_failover_policy_to_value( let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value( append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts { build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -132,7 +138,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
body_rules: transport.endpoint.body_rules.as_ref(), body_rules: transport.endpoint.body_rules.as_ref(),
provider_request_method: Some(serde_json::Value::Null), provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&provider_request_headers), provider_request_headers: Some(&provider_request_headers),
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -160,39 +166,39 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
&transport, &transport,
); );
Ok(Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream,
decision_is_stream, decision_kind: decision_kind.to_string(),
decision_kind: decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.clone(),
candidate_id: candidate_id.clone(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: None,
provider_request_method: None, auth_header: Some(auth_header),
auth_header: Some(auth_header), auth_value: Some(auth_value),
auth_value: Some(auth_value), provider_api_format,
provider_api_format, client_api_format: "openai:chat".to_string(),
client_api_format: "openai:chat".to_string(), model_name: input.requested_model.clone(),
model_name: input.requested_model.clone(), mapped_model,
mapped_model, prompt_cache_key,
prompt_cache_key, provider_request_headers,
provider_request_headers, provider_request_body: Some(provider_request_body),
provider_request_body: Some(provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: Some("application/json".to_string()),
content_type: Some("application/json".to_string()), proxy,
proxy, transport_profile,
transport_profile, timeouts,
timeouts, upstream_is_stream,
upstream_is_stream, report_kind: Some(report_kind),
report_kind: Some(report_kind), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
))) Ok(Some(decision))
} }

View File

@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value; use serde_json::Value;
use crate::ai_serving::planner::candidate_preparation::{ use crate::ai_serving::planner::candidate_preparation::{
@@ -27,8 +28,9 @@ use crate::ai_serving::transport::kiro::{
}; };
use crate::ai_serving::transport::local_openai_chat_transport_unsupported_reason; use crate::ai_serving::transport::local_openai_chat_transport_unsupported_reason;
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers, build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
StandardProviderRequestHeadersInput, build_standard_provider_request_headers, GrokHeaderInput, StandardProviderRequestHeadersInput,
GROK_CHAT_PATH,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
ai_local_execution_contract_for_formats, request_conversion_direct_auth, ai_local_execution_contract_for_formats, request_conversion_direct_auth,
@@ -64,6 +66,14 @@ pub(crate) struct LocalOpenAiChatCandidatePayloadParts {
pub(super) envelope_name: Option<&'static str>, pub(super) envelope_name: Option<&'static str>,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>, pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) request_redacted: bool, pub(super) request_redacted: bool,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
}
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
} }
fn request_identity_response_encoding_when_redacted( fn request_identity_response_encoding_when_redacted(
@@ -177,6 +187,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str(); let provider_api_format = eligible.provider_api_format.as_str();
let transport = &eligible.transport; let transport = &eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let force_body_stream_field = let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref()); endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let enable_model_directives = let enable_model_directives =
@@ -190,6 +201,130 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
resolve_provider_chat_request_redaction(state, parts, body_json, input, candidate_id) resolve_provider_chat_request_redaction(state, parts, body_json, input, candidate_id)
.await?; .await?;
let body_json = redaction.body_json.as_ref(); let body_json = redaction.body_json.as_ref();
let effective_headers = input.effective_headers(&parts.headers);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
let prepared_candidate = match prepare_header_authenticated_candidate(
planner_state,
transport,
candidate,
crate::ai_serving::transport::resolve_grok_session_auth(transport),
OauthPreparationContext {
trace_id,
api_format: provider_api_format,
operation: "openai_chat_same_format",
},
)
.await
{
Ok(prepared) => prepared,
Err(skip_reason) => {
mark_skipped_local_openai_chat_candidate(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
skip_reason,
)
.await;
return Ok(None);
}
};
let Some(provider_request_body) = build_local_openai_chat_request_body(
body_json,
&prepared_candidate.mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.endpoint.body_rules.as_ref(),
effective_headers,
enable_model_directives,
) else {
mark_skipped_local_openai_chat_candidate_with_extra_data(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"provider_request_body_build_failed",
request_body_build_failure_extra_data(
body_json,
"openai:chat",
provider_api_format,
),
)
.await;
return Ok(None);
};
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
let Some(mut provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
transport,
transport_profile: transport_profile.as_ref(),
request_headers: Some(effective_headers),
content_type: "application/json",
accept: "text/event-stream",
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
}) else {
mark_skipped_local_openai_chat_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
"openai:chat",
provider_api_format,
"grok_openai_chat_headers",
),
)
.await;
return Ok(None);
};
let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats("openai:chat", provider_api_format);
let resolved_report_kind =
if decision_kind == OPENAI_CHAT_STREAM_PLAN_KIND || !upstream_is_stream {
report_kind.to_string()
} else {
"openai_chat_sync_finalize".to_string()
};
request_identity_response_encoding_when_redacted(
&mut provider_request_headers,
redaction.redacted,
);
return Ok(Some(LocalOpenAiChatCandidatePayloadParts {
auth_header: prepared_candidate.auth_header,
auth_value: prepared_candidate.auth_value,
mapped_model: prepared_candidate.mapped_model,
provider_api_format: provider_api_format.to_string(),
provider_request_body,
provider_request_headers,
upstream_url,
execution_strategy,
conversion_mode,
report_kind: resolved_report_kind,
envelope_name: None,
transport: Arc::clone(transport),
request_redacted: redaction.redacted,
transport_profile,
}));
}
if provider_api_format == "openai:chat" { if provider_api_format == "openai:chat" {
if let Some(skip_reason) = local_openai_chat_transport_unsupported_reason(transport) { if let Some(skip_reason) = local_openai_chat_transport_unsupported_reason(transport) {
@@ -241,7 +376,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
upstream_is_stream, upstream_is_stream,
force_body_stream_field, force_body_stream_field,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
&parts.headers, effective_headers,
enable_model_directives, enable_model_directives,
) else { ) else {
mark_skipped_local_openai_chat_candidate_with_extra_data( mark_skipped_local_openai_chat_candidate_with_extra_data(
@@ -286,7 +421,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport, transport,
provider_api_format, provider_api_format,
same_format: true, same_format: true,
headers: &parts.headers, headers: effective_headers,
auth_header: &prepared_candidate.auth_header, auth_header: &prepared_candidate.auth_header,
auth_value: &prepared_candidate.auth_value, auth_value: &prepared_candidate.auth_value,
extra_headers: &BTreeMap::new(), extra_headers: &BTreeMap::new(),
@@ -317,7 +452,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
apply_codex_openai_responses_special_headers( apply_codex_openai_responses_special_headers(
&mut provider_request_headers, &mut provider_request_headers,
&provider_request_body, &provider_request_body,
&parts.headers, effective_headers,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
transport.endpoint.api_format.as_str(), transport.endpoint.api_format.as_str(),
Some(trace_id), Some(trace_id),
@@ -351,6 +486,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
envelope_name: None, envelope_name: None,
transport: Arc::clone(transport), transport: Arc::clone(transport),
request_redacted: redaction.redacted, request_redacted: redaction.redacted,
transport_profile,
})); }));
}; };
@@ -480,7 +616,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport.endpoint.body_rules.as_ref() transport.endpoint.body_rules.as_ref()
}, },
Some(input.auth_context.api_key_id.as_str()), Some(input.auth_context.api_key_id.as_str()),
&parts.headers, effective_headers,
enable_model_directives, enable_model_directives,
) else { ) else {
mark_skipped_local_openai_chat_candidate_with_extra_data( mark_skipped_local_openai_chat_candidate_with_extra_data(
@@ -575,7 +711,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
transport, transport,
provider_api_format: provider_api_format.as_str(), provider_api_format: provider_api_format.as_str(),
same_format: false, same_format: false,
headers: &parts.headers, headers: effective_headers,
auth_header: &prepared_candidate.auth_header, auth_header: &prepared_candidate.auth_header,
auth_value: &prepared_candidate.auth_value, auth_value: &prepared_candidate.auth_value,
extra_headers: &BTreeMap::new(), extra_headers: &BTreeMap::new(),
@@ -606,7 +742,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
apply_codex_openai_responses_special_headers( apply_codex_openai_responses_special_headers(
&mut provider_request_headers, &mut provider_request_headers,
&provider_request_body, &provider_request_body,
&parts.headers, effective_headers,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
provider_api_format.as_str(), provider_api_format.as_str(),
Some(trace_id), Some(trace_id),
@@ -639,6 +775,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
envelope_name: None, envelope_name: None,
transport: Arc::clone(transport), transport: Arc::clone(transport),
request_redacted: redaction.redacted, request_redacted: redaction.redacted,
transport_profile: None,
})) }))
} }
@@ -664,12 +801,13 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
request_redacted: bool, request_redacted: bool,
) -> Option<LocalOpenAiChatCandidatePayloadParts> { ) -> Option<LocalOpenAiChatCandidatePayloadParts> {
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let effective_headers = input.effective_headers(&parts.headers);
let provider_request_body = match build_kiro_provider_request_body( let provider_request_body = match build_kiro_provider_request_body(
&claude_request_body, &claude_request_body,
&mapped_model, &mapped_model,
&kiro_auth.auth_config, &kiro_auth.auth_config,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
) { ) {
Some(body) => body, Some(body) => body,
None => { None => {
@@ -720,7 +858,7 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
} }
}; };
let mut provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput { let mut provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
headers: &parts.headers, headers: effective_headers,
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: original_body_json, original_request_body: original_body_json,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
@@ -775,6 +913,7 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
envelope_name: Some(KIRO_ENVELOPE_NAME), envelope_name: Some(KIRO_ENVELOPE_NAME),
transport: Arc::clone(transport), transport: Arc::clone(transport),
request_redacted, request_redacted,
transport_profile: None,
}) })
} }

View File

@@ -140,6 +140,7 @@ pub(crate) async fn materialize_local_openai_chat_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -220,6 +221,7 @@ pub(crate) async fn build_local_openai_chat_candidate_attempt_source<'a>(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -298,6 +300,7 @@ pub(crate) async fn build_lazy_local_openai_chat_candidate_attempt_source<'a>(
&input.auth_snapshot, &input.auth_snapshot,
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,

View File

@@ -136,10 +136,11 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, false, state, parts, trace_id, decision, body_json, plan_kind, false,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source( let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, false, state, trace_id, &input, body_json, false,
@@ -187,10 +188,11 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, false, state, parts, trace_id, decision, body_json, plan_kind, false,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source( let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, true, state, trace_id, &input, body_json, true,

View File

@@ -26,6 +26,7 @@ pub(crate) async fn list_local_openai_chat_candidates(
require_streaming, require_streaming,
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
&input.auth_snapshot, &input.auth_snapshot,
input.routing_policy.as_ref(),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
false, false,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,

View File

@@ -4,11 +4,12 @@ use super::super::{GatewayControlDecision, LocalOpenAiChatDecisionInput};
use super::diagnostic::set_local_openai_chat_miss_diagnostic; use super::diagnostic::set_local_openai_chat_miss_diagnostic;
use crate::ai_serving::planner::common::extract_standard_requested_model; use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::resolve_local_decision_execution_runtime_auth_context; use crate::ai_serving::resolve_local_decision_execution_runtime_auth_context;
use crate::client_session_affinity::client_session_affinity_from_parts; use crate::client_session_affinity::client_session_affinity_from_parts;
use crate::AppState; use crate::{AppState, GatewayError};
pub(crate) async fn resolve_local_openai_chat_decision_input( pub(crate) async fn resolve_local_openai_chat_decision_input(
state: &AppState, state: &AppState,
@@ -18,7 +19,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
body_json: &serde_json::Value, body_json: &serde_json::Value,
plan_kind: &str, plan_kind: &str,
record_miss_diagnostic: bool, record_miss_diagnostic: bool,
) -> Option<LocalOpenAiChatDecisionInput> { ) -> Result<Option<LocalOpenAiChatDecisionInput>, GatewayError> {
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
@@ -37,7 +38,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
"missing_auth_context", "missing_auth_context",
); );
} }
return None; return Ok(None);
}; };
let Some(requested_model) = extract_standard_requested_model(body_json) else { let Some(requested_model) = extract_standard_requested_model(body_json) else {
@@ -55,7 +56,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
"missing_requested_model", "missing_requested_model",
); );
} }
return None; return Ok(None);
}; };
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
@@ -84,7 +85,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
"auth_snapshot_missing", "auth_snapshot_missing",
); );
} }
return None; return Ok(None);
} }
Err(err) => { Err(err) => {
warn!( warn!(
@@ -102,12 +103,28 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
"auth_snapshot_read_failed", "auth_snapshot_read_failed",
); );
} }
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
"openai:chat",
)
.await
{
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai chat decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }

View File

@@ -23,7 +23,7 @@ pub(crate) struct LocalOpenAiChatStreamAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalOpenAiChatDecisionInput, input: LocalOpenAiChatDecisionInput,
candidates: LocalOpenAiChatCandidateAttemptSource<'a>, candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
} }
@@ -43,13 +43,18 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, true, state, parts, trace_id, decision, body_json, plan_kind, true,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source( let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, true, state,
trace_id,
&input,
&effective_body_json,
true,
) )
.await; .await;
if candidate_count == 0 { if candidate_count == 0 {
@@ -77,7 +82,7 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
candidates, candidates,
}, },
@@ -127,7 +132,7 @@ impl LocalOpenAiChatStreamAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
OPENAI_CHAT_STREAM_PLAN_KIND, OPENAI_CHAT_STREAM_PLAN_KIND,
@@ -139,7 +144,7 @@ impl LocalOpenAiChatStreamAttemptSource<'_> {
return Ok(None); return Ok(None);
}; };
match build_openai_chat_stream_plan_from_decision(self.parts, self.body_json, payload) { match build_openai_chat_stream_plan_from_decision(self.parts, &self.body_json, payload) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
Err(err) => { Err(err) => {
warn!( warn!(
@@ -168,7 +173,7 @@ pub(crate) async fn build_local_openai_chat_stream_plan_and_reports(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, true, state, parts, trace_id, decision, body_json, plan_kind, true,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };

View File

@@ -23,7 +23,7 @@ pub(crate) struct LocalOpenAiChatSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalOpenAiChatDecisionInput, input: LocalOpenAiChatDecisionInput,
candidates: LocalOpenAiChatCandidateAttemptSource<'a>, candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
} }
@@ -43,13 +43,18 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, true, state, parts, trace_id, decision, body_json, plan_kind, true,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source( let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source(
state, trace_id, &input, body_json, false, state,
trace_id,
&input,
&effective_body_json,
false,
) )
.await; .await;
if candidate_count == 0 { if candidate_count == 0 {
@@ -77,7 +82,7 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
candidates, candidates,
}, },
@@ -127,7 +132,7 @@ impl LocalOpenAiChatSyncAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
OPENAI_CHAT_SYNC_PLAN_KIND, OPENAI_CHAT_SYNC_PLAN_KIND,
@@ -139,7 +144,7 @@ impl LocalOpenAiChatSyncAttemptSource<'_> {
return Ok(None); return Ok(None);
}; };
match build_openai_chat_sync_plan_from_decision(self.parts, self.body_json, payload) { match build_openai_chat_sync_plan_from_decision(self.parts, &self.body_json, payload) {
Ok(value) => Ok(value), Ok(value) => Ok(value),
Err(err) => { Err(err) => {
warn!( warn!(
@@ -168,7 +173,7 @@ pub(crate) async fn build_local_openai_chat_sync_plan_and_reports(
let Some(input) = resolve_local_openai_chat_decision_input( let Some(input) = resolve_local_openai_chat_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, true, state, parts, trace_id, decision, body_json, plan_kind, true,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };

View File

@@ -2,6 +2,7 @@ use serde_json::json;
use tracing::debug; use tracing::debug;
use crate::ai_serving::build_request_trace_proxy_value; use crate::ai_serving::build_request_trace_proxy_value;
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
use crate::ai_serving::planner::report_context::{ use crate::ai_serving::planner::report_context::{
build_local_execution_report_context, insert_provider_stream_event_api_format, build_local_execution_report_context, insert_provider_stream_event_api_format,
LocalExecutionReportContextParts, LocalExecutionReportContextParts,
@@ -15,7 +16,7 @@ use crate::ai_serving::transport::{
}; };
use crate::{ use crate::{
append_execution_contract_fields_to_value, append_local_failover_policy_to_value, append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
AiExecutionDecision, AppState, AiExecutionDecision, AppState, GatewayError,
}; };
use super::request::resolve_local_openai_responses_candidate_payload_parts; use super::request::resolve_local_openai_responses_candidate_payload_parts;
@@ -30,7 +31,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
input: &LocalOpenAiResponsesDecisionInput, input: &LocalOpenAiResponsesDecisionInput,
attempt: LocalOpenAiResponsesCandidateAttempt, attempt: LocalOpenAiResponsesCandidateAttempt,
spec: LocalOpenAiResponsesSpec, spec: LocalOpenAiResponsesSpec,
) -> Option<AiExecutionDecision> { ) -> Result<Option<AiExecutionDecision>, GatewayError> {
let spec_metadata = local_openai_responses_spec_metadata(spec); let spec_metadata = local_openai_responses_spec_metadata(spec);
let attempt_identity = attempt.attempt_identity(); let attempt_identity = attempt.attempt_identity();
let LocalOpenAiResponsesCandidateAttempt { let LocalOpenAiResponsesCandidateAttempt {
@@ -39,7 +40,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
candidate_id, candidate_id,
.. ..
} = attempt; } = attempt;
let resolved = resolve_local_openai_responses_candidate_payload_parts( let Some(resolved) = resolve_local_openai_responses_candidate_payload_parts(
state, state,
parts, parts,
trace_id, trace_id,
@@ -50,7 +51,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
&candidate_id, &candidate_id,
spec, spec,
) )
.await?; .await
else {
return Ok(None);
};
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let prompt_cache_key = resolved let prompt_cache_key = resolved
@@ -63,7 +67,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
let proxy = state let proxy = state
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport) .resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
.await; .await;
let transport_profile = resolve_transport_profile(&resolved.transport); let transport_profile = resolved
.transport_profile
.clone()
.or_else(|| resolve_transport_profile(&resolved.transport));
let timeouts = resolve_transport_execution_timeouts(&resolved.transport); let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let mut extra_fields = serde_json::Map::new(); let mut extra_fields = serde_json::Map::new();
if let Some(proxy_value) = if let Some(proxy_value) =
@@ -78,6 +85,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
&mut extra_fields, &mut extra_fields,
resolved.transport.provider.provider_type.as_str(), resolved.transport.provider.provider_type.as_str(),
); );
let effective_headers = input.effective_headers(&parts.headers);
let report_context = append_local_failover_policy_to_value( let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value( append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts { build_local_execution_report_context(LocalExecutionReportContextParts {
@@ -105,7 +113,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
body_rules: resolved.transport.endpoint.body_rules.as_ref(), body_rules: resolved.transport.endpoint.body_rules.as_ref(),
provider_request_method: Some(serde_json::Value::Null), provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers), provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers, original_headers: effective_headers,
request_path: Some(parts.uri.path()), request_path: Some(parts.uri.path()),
request_query_string: parts.uri.query(), request_query_string: parts.uri.query(),
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)), request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
@@ -170,41 +178,42 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
envelope_name: _, envelope_name: _,
upstream_is_stream, upstream_is_stream,
transport, transport,
transport_profile: _,
} = resolved; } = resolved;
Some(build_ai_execution_decision_response( let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
AiExecutionDecisionResponseParts { decision_is_stream: spec_metadata.require_streaming,
decision_is_stream: spec_metadata.require_streaming, decision_kind: spec_metadata.decision_kind.to_string(),
decision_kind: spec_metadata.decision_kind.to_string(), execution_strategy,
execution_strategy, conversion_mode,
conversion_mode, request_id: trace_id.to_string(),
request_id: trace_id.to_string(), candidate_id: candidate_id.clone(),
candidate_id: candidate_id.clone(), provider_name: transport.provider.name.clone(),
provider_name: transport.provider.name.clone(), provider_id: candidate.provider_id.clone(),
provider_id: candidate.provider_id.clone(), endpoint_id: candidate.endpoint_id.clone(),
endpoint_id: candidate.endpoint_id.clone(), key_id: candidate.key_id.clone(),
key_id: candidate.key_id.clone(), upstream_base_url: transport.endpoint.base_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(), upstream_url,
upstream_url, provider_request_method: None,
provider_request_method: None, auth_header: Some(auth_header),
auth_header: Some(auth_header), auth_value: Some(auth_value),
auth_value: Some(auth_value), provider_api_format,
provider_api_format, client_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(), model_name: input.requested_model.clone(),
model_name: input.requested_model.clone(), mapped_model,
mapped_model, prompt_cache_key,
prompt_cache_key, provider_request_headers,
provider_request_headers, provider_request_body: Some(provider_request_body),
provider_request_body: Some(provider_request_body), provider_request_body_base64: None,
provider_request_body_base64: None, content_type: Some("application/json".to_string()),
content_type: Some("application/json".to_string()), proxy,
proxy, transport_profile,
transport_profile, timeouts,
timeouts, upstream_is_stream,
upstream_is_stream, report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned), report_context: Some(report_context),
report_context: Some(report_context), auth_context: input.auth_context.clone(),
auth_context: input.auth_context.clone(), });
}, apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
)) Ok(Some(decision))
} }

View File

@@ -1,6 +1,7 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::sync::Arc; use std::sync::Arc;
use aether_contracts::ResolvedTransportProfile;
use serde_json::Value; use serde_json::Value;
use tracing::debug; use tracing::debug;
@@ -35,8 +36,10 @@ use crate::ai_serving::transport::kiro::{
KiroRequestAuth, KIRO_ENVELOPE_NAME, KiroRequestAuth, KIRO_ENVELOPE_NAME,
}; };
use crate::ai_serving::transport::{ use crate::ai_serving::transport::{
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers, build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
local_standard_transport_unsupported_reason_with_network, StandardProviderRequestHeadersInput, build_standard_provider_request_headers,
local_standard_transport_unsupported_reason_with_network, GrokHeaderInput,
StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
ai_local_execution_contract_for_formats, request_conversion_direct_auth, ai_local_execution_contract_for_formats, request_conversion_direct_auth,
@@ -56,6 +59,13 @@ use super::LocalOpenAiResponsesSpec;
const ANTIGRAVITY_ENVELOPE_NAME: &str = "antigravity:v1internal"; const ANTIGRAVITY_ENVELOPE_NAME: &str = "antigravity:v1internal";
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
matches!(
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
)
}
pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts { pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
pub(super) auth_header: String, pub(super) auth_header: String,
pub(super) auth_value: String, pub(super) auth_value: String,
@@ -70,6 +80,7 @@ pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
pub(super) envelope_name: Option<&'static str>, pub(super) envelope_name: Option<&'static str>,
pub(super) upstream_is_stream: bool, pub(super) upstream_is_stream: bool,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>, pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) transport_profile: Option<ResolvedTransportProfile>,
} }
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
@@ -90,12 +101,22 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let provider_api_format = eligible.provider_api_format.as_str(); let provider_api_format = eligible.provider_api_format.as_str();
let transport = &eligible.transport; let transport = &eligible.transport;
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
let is_antigravity = is_antigravity_provider_transport(transport); let is_antigravity = is_antigravity_provider_transport(transport);
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format); let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
let is_grok = transport
.provider
.provider_type
.trim()
.eq_ignore_ascii_case("grok");
let same_format = api_format_alias_matches(provider_api_format, &client_api_format); let same_format = api_format_alias_matches(provider_api_format, &client_api_format);
let conversion_kind = request_conversion_kind(spec_metadata.api_format, provider_api_format); let conversion_kind = request_conversion_kind(spec_metadata.api_format, provider_api_format);
let transport_unsupported_reason = if same_format && is_kiro_claude_cli { let transport_unsupported_reason = if is_grok
&& is_grok_text_provider_api_format(provider_api_format)
{
None
} else if same_format && is_kiro_claude_cli {
local_kiro_request_transport_unsupported_reason_with_network(transport) local_kiro_request_transport_unsupported_reason_with_network(transport)
} else if same_format { } else if same_format {
local_standard_transport_unsupported_reason_with_network(transport, provider_api_format) local_standard_transport_unsupported_reason_with_network(transport, provider_api_format)
@@ -154,7 +175,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
None None
}; };
let direct_auth = if kiro_auth.is_some() { let direct_auth = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
crate::ai_serving::transport::resolve_grok_session_auth(transport)
} else if kiro_auth.is_some() {
None None
} else if same_format { } else if same_format {
match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() { match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() {
@@ -236,42 +259,58 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
); );
let force_body_stream_field = let force_body_stream_field =
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref()); endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
let Some(mut base_provider_request_body) = (if needs_bidirectional_conversion { let effective_headers = input.effective_headers(&parts.headers);
build_cross_format_openai_responses_request_body( let Some(mut base_provider_request_body) =
body_json, (if is_grok && is_grok_text_provider_api_format(provider_api_format) {
&mapped_model, build_local_openai_responses_request_body(
spec_metadata.api_format, body_json,
provider_api_format, &mapped_model,
upstream_is_stream, upstream_is_stream,
force_body_stream_field, force_body_stream_field,
transport.provider.provider_type.as_str(), transport.provider.provider_type.as_str(),
if is_kiro_claude_cli { spec_metadata.api_format,
None transport.endpoint.body_rules.as_ref(),
} else { Some(input.auth_context.api_key_id.as_str()),
transport.endpoint.body_rules.as_ref() effective_headers,
}, enable_model_directives,
Some(input.auth_context.api_key_id.as_str()), )
&parts.headers, } else if needs_bidirectional_conversion {
enable_model_directives, build_cross_format_openai_responses_request_body(
) body_json,
} else { &mapped_model,
build_local_openai_responses_request_body( spec_metadata.api_format,
body_json, provider_api_format,
&mapped_model, upstream_is_stream,
upstream_is_stream, force_body_stream_field,
force_body_stream_field, transport.provider.provider_type.as_str(),
transport.provider.provider_type.as_str(), if is_kiro_claude_cli {
provider_api_format, None
if is_kiro_claude_cli { } else {
None transport.endpoint.body_rules.as_ref()
} else { },
transport.endpoint.body_rules.as_ref() Some(input.auth_context.api_key_id.as_str()),
}, effective_headers,
Some(input.auth_context.api_key_id.as_str()), enable_model_directives,
&parts.headers, )
enable_model_directives, } else {
) build_local_openai_responses_request_body(
}) else { body_json,
&mapped_model,
upstream_is_stream,
force_body_stream_field,
transport.provider.provider_type.as_str(),
provider_api_format,
if is_kiro_claude_cli {
None
} else {
transport.endpoint.body_rules.as_ref()
},
Some(input.auth_context.api_key_id.as_str()),
effective_headers,
enable_model_directives,
)
})
else {
mark_skipped_local_openai_responses_candidate_with_extra_data( mark_skipped_local_openai_responses_candidate_with_extra_data(
state, state,
input, input,
@@ -390,7 +429,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
.await; .await;
} }
let Some(upstream_url) = (if needs_bidirectional_conversion { let Some(upstream_url) = (if is_grok && is_grok_text_provider_api_format(provider_api_format) {
Some(build_grok_upstream_url(transport, GROK_CHAT_PATH))
} else if needs_bidirectional_conversion {
build_cross_format_openai_responses_upstream_url( build_cross_format_openai_responses_upstream_url(
parts, parts,
transport, transport,
@@ -427,48 +468,86 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
.as_ref() .as_ref()
.map(build_antigravity_static_identity_headers) .map(build_antigravity_static_identity_headers)
.unwrap_or_default(); .unwrap_or_default();
let Some(resolved_headers) = let resolved_headers = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
build_standard_provider_request_headers(StandardProviderRequestHeadersInput { let Some(headers) = build_grok_browser_headers(GrokHeaderInput {
transport, transport,
provider_api_format, transport_profile: transport_profile.as_ref(),
same_format, request_headers: Some(effective_headers),
headers: &parts.headers, content_type: "application/json",
auth_header: &auth_header, accept: "text/event-stream",
auth_value: &auth_value,
extra_headers: &extra_headers,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: body_json, original_request_body: body_json,
upstream_is_stream, }) else {
}) mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
else { state,
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic( input,
state, trace_id,
input, candidate,
trace_id, candidate_index,
candidate, candidate_id,
candidate_index, "transport_header_rules_apply_failed",
candidate_id, CandidateFailureDiagnostic::header_rules_apply_failed(
"transport_header_rules_apply_failed", spec_metadata.api_format,
CandidateFailureDiagnostic::header_rules_apply_failed( provider_api_format,
spec_metadata.api_format, "grok_openai_responses_headers",
),
)
.await;
return None;
};
crate::ai_serving::transport::StandardProviderRequestHeaders {
headers,
auth_header: auth_header.clone(),
auth_value: auth_value.clone(),
}
} else {
let Some(resolved_headers) =
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
transport,
provider_api_format, provider_api_format,
"openai_responses_headers", same_format,
), headers: effective_headers,
) auth_header: &auth_header,
.await; auth_value: &auth_value,
return None; extra_headers: &extra_headers,
header_rules: transport.endpoint.header_rules.as_ref(),
provider_request_body: &provider_request_body,
original_request_body: body_json,
upstream_is_stream,
})
else {
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
state,
input,
trace_id,
candidate,
candidate_index,
candidate_id,
"transport_header_rules_apply_failed",
CandidateFailureDiagnostic::header_rules_apply_failed(
spec_metadata.api_format,
provider_api_format,
"openai_responses_headers",
),
)
.await;
return None;
};
resolved_headers
}; };
let mut provider_request_headers = resolved_headers.headers; let mut provider_request_headers = resolved_headers.headers;
apply_codex_openai_responses_special_headers( if !is_grok {
&mut provider_request_headers, apply_codex_openai_responses_special_headers(
&provider_request_body, &mut provider_request_headers,
&parts.headers, &provider_request_body,
transport.provider.provider_type.as_str(), effective_headers,
provider_api_format, transport.provider.provider_type.as_str(),
Some(trace_id), provider_api_format,
transport.key.decrypted_auth_config.as_deref(), Some(trace_id),
); transport.key.decrypted_auth_config.as_deref(),
);
}
let (execution_strategy, conversion_mode) = let (execution_strategy, conversion_mode) =
ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format); ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format);
@@ -516,6 +595,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
}, },
upstream_is_stream, upstream_is_stream,
transport: Arc::clone(transport), transport: Arc::clone(transport),
transport_profile,
}) })
} }
@@ -545,12 +625,13 @@ async fn build_kiro_openai_responses_payload_parts(
kiro_auth: &KiroRequestAuth, kiro_auth: &KiroRequestAuth,
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> { ) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
let candidate = &eligible.candidate; let candidate = &eligible.candidate;
let effective_headers = input.effective_headers(&parts.headers);
let provider_request_body = match build_kiro_provider_request_body( let provider_request_body = match build_kiro_provider_request_body(
&claude_request_body, &claude_request_body,
&mapped_model, &mapped_model,
&kiro_auth.auth_config, &kiro_auth.auth_config,
transport.endpoint.body_rules.as_ref(), transport.endpoint.body_rules.as_ref(),
Some(&parts.headers), Some(effective_headers),
) { ) {
Some(body) => body, Some(body) => body,
None => { None => {
@@ -601,7 +682,7 @@ async fn build_kiro_openai_responses_payload_parts(
} }
}; };
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput { let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
headers: &parts.headers, headers: effective_headers,
provider_request_body: &provider_request_body, provider_request_body: &provider_request_body,
original_request_body: original_body_json, original_request_body: original_body_json,
header_rules: transport.endpoint.header_rules.as_ref(), header_rules: transport.endpoint.header_rules.as_ref(),
@@ -666,5 +747,6 @@ async fn build_kiro_openai_responses_payload_parts(
envelope_name: Some(KIRO_ENVELOPE_NAME), envelope_name: Some(KIRO_ENVELOPE_NAME),
upstream_is_stream, upstream_is_stream,
transport: Arc::clone(transport), transport: Arc::clone(transport),
transport_profile: None,
}) })
} }

View File

@@ -19,6 +19,7 @@ use crate::ai_serving::planner::candidate_source::{
}; };
use crate::ai_serving::planner::common::extract_standard_requested_model; use crate::ai_serving::planner::common::extract_standard_requested_model;
use crate::ai_serving::planner::decision_input::{ use crate::ai_serving::planner::decision_input::{
attach_routing_policy_to_local_requested_model_input,
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input, build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
}; };
use crate::ai_serving::planner::materialization_policy::{ use crate::ai_serving::planner::materialization_policy::{
@@ -48,7 +49,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
decision: &GatewayControlDecision, decision: &GatewayControlDecision,
body_json: &serde_json::Value, body_json: &serde_json::Value,
plan_kind: &str, plan_kind: &str,
) -> Option<LocalOpenAiResponsesDecisionInput> { ) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else { let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
warn!( warn!(
trace_id = %trace_id, trace_id = %trace_id,
@@ -65,7 +66,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
extract_standard_requested_model(body_json).as_deref(), extract_standard_requested_model(body_json).as_deref(),
"missing_auth_context", "missing_auth_context",
); );
return None; return Ok(None);
}; };
let Some(requested_model) = extract_standard_requested_model(body_json) else { let Some(requested_model) = extract_standard_requested_model(body_json) else {
@@ -81,7 +82,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
None, None,
"missing_requested_model", "missing_requested_model",
); );
return None; return Ok(None);
}; };
let resolved_input = match resolve_local_authenticated_decision_input( let resolved_input = match resolve_local_authenticated_decision_input(
@@ -108,7 +109,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
Some(requested_model.as_str()), Some(requested_model.as_str()),
"auth_snapshot_missing", "auth_snapshot_missing",
); );
return None; return Ok(None);
} }
Err(err) => { Err(err) => {
warn!( warn!(
@@ -124,14 +125,30 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
Some(requested_model.as_str()), Some(requested_model.as_str()),
"auth_snapshot_read_failed", "auth_snapshot_read_failed",
); );
return None; return Err(err);
} }
}; };
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model); let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
input.request_auth_channel = decision.request_auth_channel.clone(); input.request_auth_channel = decision.request_auth_channel.clone();
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json)); input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
Some(input) if let Err(err) = attach_routing_policy_to_local_requested_model_input(
state,
parts,
&mut input,
body_json,
"openai:responses",
)
.await
{
warn!(
trace_id = %trace_id,
error = ?err,
"gateway local openai responses decision routing profile resolution failed"
);
return Err(err);
}
Ok(Some(input))
} }
pub(crate) async fn materialize_local_openai_responses_candidate_attempts( pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
@@ -157,6 +174,7 @@ pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
spec_metadata.require_streaming, spec_metadata.require_streaming,
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
&input.auth_snapshot, &input.auth_snapshot,
input.routing_policy.as_ref(),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
true, true,
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat, LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
@@ -170,6 +188,7 @@ pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
Some(&input.auth_snapshot), Some(&input.auth_snapshot),
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,
@@ -256,6 +275,7 @@ pub(crate) async fn build_local_openai_responses_candidate_attempt_source<'a>(
&input.auth_snapshot, &input.auth_snapshot,
input.client_session_affinity.as_ref(), input.client_session_affinity.as_ref(),
input.required_capabilities.as_ref(), input.required_capabilities.as_ref(),
input.routing_policy.as_ref(),
sticky_session_token.as_deref(), sticky_session_token.as_deref(),
input.request_auth_channel.as_deref(), input.request_auth_channel.as_deref(),
persistence_policy, persistence_policy,

View File

@@ -103,10 +103,11 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
let Some(input) = resolve_local_openai_responses_decision_input( let Some(input) = resolve_local_openai_responses_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, state, parts, trace_id, decision, body_json, plan_kind,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_local_openai_responses_candidate_attempt_source( let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
@@ -117,7 +118,7 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate( if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }
@@ -141,10 +142,11 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
let Some(input) = resolve_local_openai_responses_decision_input( let Some(input) = resolve_local_openai_responses_decision_input(
state, parts, trace_id, decision, body_json, plan_kind, state, parts, trace_id, decision, body_json, plan_kind,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
let body_json = input.effective_body_json(body_json);
let (mut source, _) = build_local_openai_responses_candidate_attempt_source( let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state, trace_id, &input, body_json, spec,
@@ -155,7 +157,7 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate( if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
{ {
return Ok(Some(payload)); return Ok(Some(payload));
} }

View File

@@ -29,7 +29,7 @@ pub(crate) struct LocalOpenAiResponsesSyncAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalOpenAiResponsesDecisionInput, input: LocalOpenAiResponsesDecisionInput,
spec: LocalOpenAiResponsesSpec, spec: LocalOpenAiResponsesSpec,
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>, candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
@@ -39,7 +39,7 @@ pub(crate) struct LocalOpenAiResponsesStreamAttemptSource<'a> {
state: &'a AppState, state: &'a AppState,
parts: &'a http::request::Parts, parts: &'a http::request::Parts,
trace_id: &'a str, trace_id: &'a str,
body_json: &'a serde_json::Value, body_json: serde_json::Value,
input: LocalOpenAiResponsesDecisionInput, input: LocalOpenAiResponsesDecisionInput,
spec: LocalOpenAiResponsesSpec, spec: LocalOpenAiResponsesSpec,
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>, candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
@@ -62,7 +62,7 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
body_json, body_json,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -74,8 +74,13 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source( let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state,
trace_id,
&input,
&effective_body_json,
spec,
) )
.await?; .await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count); apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
@@ -88,7 +93,7 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
candidates, candidates,
@@ -114,7 +119,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
body_json, body_json,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
@@ -126,8 +131,13 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
Some(input.requested_model.as_str()), Some(input.requested_model.as_str()),
"candidate_evaluation_incomplete", "candidate_evaluation_incomplete",
); );
let effective_body_json = input.effective_body_json(body_json).clone();
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source( let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
state, trace_id, &input, body_json, spec, state,
trace_id,
&input,
&effective_body_json,
spec,
) )
.await?; .await?;
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count); apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
@@ -140,7 +150,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
state, state,
parts, parts,
trace_id, trace_id,
body_json, body_json: effective_body_json,
input, input,
spec, spec,
candidates, candidates,
@@ -214,19 +224,19 @@ impl LocalOpenAiResponsesSyncAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
match build_openai_responses_sync_plan_from_decision( match build_openai_responses_sync_plan_from_decision(
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
self.spec.compact, self.spec.compact,
) { ) {
@@ -252,19 +262,19 @@ impl LocalOpenAiResponsesStreamAttemptSource<'_> {
self.state, self.state,
self.parts, self.parts,
self.trace_id, self.trace_id,
self.body_json, &self.body_json,
&self.input, &self.input,
attempt, attempt,
self.spec, self.spec,
) )
.await .await?
else { else {
return Ok(None); return Ok(None);
}; };
match build_openai_responses_stream_plan_from_decision( match build_openai_responses_stream_plan_from_decision(
self.parts, self.parts,
self.body_json, &self.body_json,
payload, payload,
self.spec.compact, self.spec.compact,
) { ) {
@@ -298,7 +308,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
body_json, body_json,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
@@ -325,7 +335,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate( let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };
@@ -370,7 +380,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
body_json, body_json,
spec_metadata.decision_kind, spec_metadata.decision_kind,
) )
.await .await?
else { else {
return Ok(Vec::new()); return Ok(Vec::new());
}; };
@@ -397,7 +407,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate( let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
state, parts, trace_id, body_json, &input, attempt, spec, state, parts, trace_id, body_json, &input, attempt, spec,
) )
.await .await?
else { else {
continue; continue;
}; };

View File

@@ -61,6 +61,7 @@ pub(crate) use aether_ai_formats::api::{
maybe_build_standard_sync_finalize_product_from_normalized_payload, model_directive_base_model, maybe_build_standard_sync_finalize_product_from_normalized_payload, model_directive_base_model,
normalize_api_format_alias, normalize_claude_request_to_openai_chat_request, normalize_api_format_alias, normalize_claude_request_to_openai_chat_request,
normalize_gemini_request_to_openai_chat_request, normalize_openai_image_request, normalize_gemini_request_to_openai_chat_request, normalize_openai_image_request,
normalize_openai_image_request_with_options,
normalize_openai_responses_request_to_openai_chat_request, normalize_openai_responses_request_to_openai_chat_request,
normalize_provider_private_report_context, normalize_provider_private_response_value, normalize_provider_private_report_context, normalize_provider_private_response_value,
normalize_standard_request_to_openai_chat_request, openai_image_operation_from_path, normalize_standard_request_to_openai_chat_request, openai_image_operation_from_path,
@@ -97,10 +98,10 @@ pub(crate) use aether_ai_formats::api::{
LocalStandardSourceMode, LocalStandardSpec, LocalSyncReportParts, LocalVideoCreateFamily, LocalStandardSourceMode, LocalStandardSpec, LocalSyncReportParts, LocalVideoCreateFamily,
LocalVideoCreateSpec, NormalizedOpenAiImageRequest, OpenAIChatClientEmitter, LocalVideoCreateSpec, NormalizedOpenAiImageRequest, OpenAIChatClientEmitter,
OpenAIChatProviderState, OpenAIResponsesClientEmitter, OpenAIResponsesProviderState, OpenAIChatProviderState, OpenAIResponsesClientEmitter, OpenAIResponsesProviderState,
OpenAiImageOperation, OpenAiImageRequestForGemini, OpenAiImageResponseFormat, OpenAiImageNormalizeOptions, OpenAiImageOperation, OpenAiImageRequestForGemini,
OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct, ProviderAdaptationDescriptor, OpenAiImageResponseFormat, OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct,
ProviderAdaptationSurface, ProviderPrivateStreamNormalizer, RequestConversionKind, ProviderAdaptationDescriptor, ProviderAdaptationSurface, ProviderPrivateStreamNormalizer,
StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct, RequestConversionKind, StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct,
StreamingStandardFormatMatrix, SyncChatResponseConversionKind, SyncCliResponseConversionKind, StreamingStandardFormatMatrix, SyncChatResponseConversionKind, SyncCliResponseConversionKind,
SyncToStreamBridgeOutcome, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, CLAUDE_CHAT_STREAM_PLAN_KIND, SyncToStreamBridgeOutcome, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, CLAUDE_CHAT_STREAM_PLAN_KIND,
CLAUDE_CHAT_STREAM_SUCCESS_REPORT_KIND, CLAUDE_CHAT_SYNC_ERROR_REPORT_KIND, CLAUDE_CHAT_STREAM_SUCCESS_REPORT_KIND, CLAUDE_CHAT_SYNC_ERROR_REPORT_KIND,

View File

@@ -14,6 +14,10 @@ pub(crate) mod kiro {
pub(crate) use aether_provider_transport::kiro::*; pub(crate) use aether_provider_transport::kiro::*;
} }
pub(crate) mod grok {
pub(crate) use aether_provider_transport::grok::*;
}
pub(crate) mod oauth_refresh { pub(crate) mod oauth_refresh {
pub(crate) use aether_provider_transport::oauth_refresh::*; pub(crate) use aether_provider_transport::oauth_refresh::*;
} }
@@ -55,6 +59,7 @@ pub(crate) use aether_provider_transport::{
body_rules_handle_path, body_rules_have_enabled_rules, body_rules_handle_path, body_rules_have_enabled_rules,
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url, build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url, build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url,
build_grok_app_chat_body, build_grok_browser_headers, build_grok_upstream_url,
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url, build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
build_local_openai_responses_upstream_url, build_openai_image_headers, build_local_openai_responses_upstream_url, build_openai_image_headers,
build_openai_image_upstream_url, build_passthrough_headers, build_request_trace_proxy_value, build_openai_image_upstream_url, build_passthrough_headers, build_request_trace_proxy_value,
@@ -74,7 +79,7 @@ pub(crate) use aether_provider_transport::{
request_conversion_enabled_for_transport, request_conversion_transport_supported, request_conversion_enabled_for_transport, request_conversion_transport_supported,
request_conversion_transport_unsupported_reason, request_pair_allowed_for_transport, request_conversion_transport_unsupported_reason, request_pair_allowed_for_transport,
request_pair_direct_auth, request_pair_transport_unsupported_reason, resolve_gemini_files_auth, request_pair_direct_auth, request_pair_transport_unsupported_reason, resolve_gemini_files_auth,
resolve_openai_image_auth, resolve_same_format_provider_direct_auth, resolve_grok_session_auth, resolve_openai_image_auth, resolve_same_format_provider_direct_auth,
resolve_transport_execution_timeouts, resolve_transport_profile, resolve_transport_execution_timeouts, resolve_transport_profile,
resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity, resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
resolve_video_create_auth, same_format_provider_transport_supported, resolve_video_create_auth, same_format_provider_transport_supported,
@@ -84,12 +89,12 @@ pub(crate) use aether_provider_transport::{
supports_local_oauth_request_auth_resolution, transport_proxy_is_locally_supported, supports_local_oauth_request_auth_resolution, transport_proxy_is_locally_supported,
video_create_transport_unsupported_reason, CandidateTransportPolicyFacts, video_create_transport_unsupported_reason, CandidateTransportPolicyFacts,
GatewayProviderTransportSnapshot, GeminiFilesHeadersInput, GeminiFilesRequestBodyError, GatewayProviderTransportSnapshot, GeminiFilesHeadersInput, GeminiFilesRequestBodyError,
GeminiFilesRequestBodyParts, LocalResolvedOAuthRequestAuth, ProviderOpenAiImageHeadersInput, GeminiFilesRequestBodyParts, GrokHeaderInput, LocalResolvedOAuthRequestAuth,
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, SameFormatProviderFamily, ProviderOpenAiImageHeadersInput, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior, SameFormatProviderFamily, SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput, SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
SameFormatProviderUpstreamUrlParams, StandardPlanFallbackAcceptPolicy, SameFormatProviderUpstreamUrlParams, StandardPlanFallbackAcceptPolicy,
StandardPlanFallbackHeadersInput, StandardProviderRequestHeaders, StandardPlanFallbackHeadersInput, StandardProviderRequestHeaders,
StandardProviderRequestHeadersInput, TransportRequestBodySemanticsError, StandardProviderRequestHeadersInput, TransportRequestBodySemanticsError,
TransportRequestUrlParams, TransportRequestUrlParams, GROK_CHAT_PATH, GROK_INTERNAL_HEADER, GROK_RATE_LIMITS_PATH,
}; };

View File

@@ -17,7 +17,6 @@ const AI_POST_ROUTE_PATTERNS: &[&str] = &[
"/v1/responses/compact", "/v1/responses/compact",
"/v1/images/generations", "/v1/images/generations",
"/v1/images/edits", "/v1/images/edits",
"/v1/images/variations",
]; ];
const AI_ANY_ROUTE_PATTERNS: &[&str] = &[ const AI_ANY_ROUTE_PATTERNS: &[&str] = &[

View File

@@ -115,7 +115,6 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
"/v1/rerank", "/v1/rerank",
"/v1/images/generations", "/v1/images/generations",
"/v1/images/edits", "/v1/images/edits",
"/v1/images/variations",
"/v1/messages", "/v1/messages",
"/v1/messages/count_tokens", "/v1/messages/count_tokens",
"/v1/responses", "/v1/responses",

View File

@@ -106,7 +106,7 @@ async fn balance_capacity_rejection(
requested_model: Option<&str>, requested_model: Option<&str>,
body: &Bytes, body: &Bytes,
) -> Result<Option<GatewayLocalAuthRejection>, GatewayError> { ) -> Result<Option<GatewayLocalAuthRejection>, GatewayError> {
if auth_context.api_key_is_standalone || auth_context.admin_bypass_limits { if auth_context.api_key_is_standalone {
return Ok(None); return Ok(None);
} }
if auth_context.local_rejection.is_some() { if auth_context.local_rejection.is_some() {
@@ -816,6 +816,43 @@ mod tests {
} }
} }
#[tokio::test]
async fn admin_bypass_limits_does_not_skip_exhausted_daily_quota_capacity() {
let context = billing_context_with_pricing(
Some(json!({
"tiers": [{
"up_to": null,
"input_price_per_1m": 1.0,
"output_price_per_1m": 2.0
}]
})),
None,
None,
None,
);
let state = state_with_quota_and_wallet(quota_availability(0.0, false), context);
let mut decision = decision_with_allowed_models(vec!["gpt-5".to_string()]);
if let Some(auth_context) = decision.auth_context.as_mut() {
auth_context.admin_bypass_limits = true;
}
let uri: Uri = "/v1/chat/completions".parse().expect("uri should parse");
let body = Bytes::from_static(
br#"{"model":"gpt-5","messages":[{"role":"user","content":"hi"}],"stream":true}"#,
);
let rejection =
request_model_local_rejection(&state, Some(&decision), &uri, &json_headers(), &body)
.await
.expect("quota rejection should resolve");
assert_eq!(
rejection,
Some(GatewayLocalAuthRejection::BalanceDenied {
remaining: Some(0.0),
})
);
}
#[tokio::test] #[tokio::test]
async fn positive_balance_still_denies_known_cost_above_available_capacity() { async fn positive_balance_still_denies_known_cost_above_available_capacity() {
let context = billing_context_with_pricing( let context = billing_context_with_pricing(

View File

@@ -9,7 +9,8 @@ pub(crate) use gate::{
request_model_local_rejection, should_buffer_request_for_local_auth, request_model_local_rejection, should_buffer_request_for_local_auth,
trusted_auth_local_rejection, GatewayLocalAuthRejection, trusted_auth_local_rejection, GatewayLocalAuthRejection,
}; };
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
pub(crate) use resolution::{ pub(crate) use resolution::{
resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext, GatewayControlAuthContext, refresh_execution_runtime_auth_context, resolve_execution_runtime_auth_context,
GatewayAdminPrincipalContext, GatewayControlAuthContext,
}; };
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};

View File

@@ -433,7 +433,14 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
let _ = trace_id; let _ = trace_id;
if let Some(auth_context) = decision.auth_context.clone() { if let Some(auth_context) = decision.auth_context.clone() {
return Ok(Some(auth_context)); return Ok(Some(
refresh_execution_runtime_auth_context(
state,
auth_context,
decision.auth_endpoint_signature.as_deref(),
)
.await?,
));
} }
let Some(auth_endpoint_signature) = decision.auth_endpoint_signature.as_deref() else { let Some(auth_endpoint_signature) = decision.auth_endpoint_signature.as_deref() else {
@@ -445,7 +452,14 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
}; };
if let Some(auth_context) = get_cached_auth_context(state, &cache_key) { if let Some(auth_context) = get_cached_auth_context(state, &cache_key) {
return Ok(Some(auth_context)); let refreshed = refresh_execution_runtime_auth_context(
state,
auth_context,
Some(auth_endpoint_signature),
)
.await?;
put_cached_auth_context(state, cache_key, refreshed.clone());
return Ok(Some(refreshed));
} }
if let Some(auth_context) = if let Some(auth_context) =
@@ -461,6 +475,56 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
Ok(None) Ok(None)
} }
pub(crate) async fn refresh_execution_runtime_auth_context(
state: &AppState,
auth_context: GatewayControlAuthContext,
auth_endpoint_signature: Option<&str>,
) -> Result<GatewayControlAuthContext, GatewayError> {
if auth_context.local_rejection.is_some() || !auth_context.access_allowed {
return Ok(auth_context);
}
let Some(auth_endpoint_signature) = auth_endpoint_signature
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return Ok(auth_context);
};
if !state.has_auth_api_key_reader()
|| auth_context.user_id.trim().is_empty()
|| auth_context.api_key_id.trim().is_empty()
{
return Ok(auth_context);
}
let snapshot = state
.data
.read_auth_api_key_snapshot(
&auth_context.user_id,
&auth_context.api_key_id,
current_unix_secs(),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let Some(snapshot) = snapshot else {
let mut denied = auth_context;
denied.access_allowed = false;
denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
denied.balance_remaining = None;
return Ok(denied);
};
let wallet_access = resolve_wallet_auth_gate(state, &snapshot).await?;
Ok(build_data_backed_auth_context(
state,
snapshot,
auth_endpoint_signature,
Some(true),
auth_context.balance_remaining,
wallet_access,
)
.await)
}
fn put_cached_auth_context( fn put_cached_auth_context(
state: &AppState, state: &AppState,
cache_key: String, cache_key: String,
@@ -609,7 +673,7 @@ async fn build_data_backed_auth_context(
.api_key_expires_at_unix_secs .api_key_expires_at_unix_secs
.is_some_and(|expires_at| expires_at < current_unix_secs()); .is_some_and(|expires_at| expires_at < current_unix_secs());
let locked_api_key = snapshot.api_key_is_locked && !snapshot.api_key_is_standalone; let locked_api_key = snapshot.api_key_is_locked && !snapshot.api_key_is_standalone;
let access_allowed = header_access_allowed let key_access_allowed = header_access_allowed
.map(|value| value && snapshot.currently_usable) .map(|value| value && snapshot.currently_usable)
.unwrap_or(snapshot.currently_usable); .unwrap_or(snapshot.currently_usable);
let wallet_remaining = wallet_access let wallet_remaining = wallet_access
@@ -656,7 +720,7 @@ async fn build_data_backed_auth_context(
user_id: snapshot.user_id, user_id: snapshot.user_id,
api_key_id: snapshot.api_key_id, api_key_id: snapshot.api_key_id,
balance_remaining: wallet_remaining.or(balance_remaining), balance_remaining: wallet_remaining.or(balance_remaining),
access_allowed, access_allowed: key_access_allowed && local_rejection.is_none(),
user_rate_limit: snapshot.user_rate_limit, user_rate_limit: snapshot.user_rate_limit,
api_key_rate_limit: snapshot.api_key_rate_limit, api_key_rate_limit: snapshot.api_key_rate_limit,
api_key_is_standalone: snapshot.api_key_is_standalone, api_key_is_standalone: snapshot.api_key_is_standalone,
@@ -835,13 +899,20 @@ mod tests {
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot, InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
}; };
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::wallet::{
InMemoryWalletRepository, StoredWalletSnapshot, WalletReadRepository,
};
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
}; };
use axum::http::{HeaderMap, Uri}; use axum::http::{HeaderMap, Uri};
use super::{resolve_data_backed_auth_context, GatewayLocalAuthRejection}; use super::{
resolve_data_backed_auth_context, resolve_execution_runtime_auth_context,
GatewayLocalAuthRejection,
};
use crate::control::auth::credentials::hash_api_key; use crate::control::auth::credentials::hash_api_key;
use crate::control::GatewayControlDecision;
use crate::data::GatewayDataState; use crate::data::GatewayDataState;
use crate::AppState; use crate::AppState;
@@ -946,6 +1017,154 @@ mod tests {
assert_eq!(repository.touch_count("key-1"), 1); assert_eq!(repository.touch_count("key-1"), 1);
} }
#[tokio::test]
async fn data_backed_auth_context_marks_wallet_denial_as_not_allowed() {
let api_key = "sk-test-empty-wallet";
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-empty-wallet", "user-empty-wallet"),
)]));
let wallet_repository = Arc::new(InMemoryWalletRepository::seed(vec![
StoredWalletSnapshot::new(
"wallet-empty".to_string(),
Some("user-empty-wallet".to_string()),
None,
0.0,
0.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
0.0,
0.0,
0.0,
0.0,
100,
)
.expect("wallet should build"),
]));
let data =
GatewayDataState::with_auth_and_wallet_for_tests(auth_repository, wallet_repository);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let mut headers = HeaderMap::new();
headers.insert(
http::header::AUTHORIZATION,
format!("Bearer {api_key}").parse().unwrap(),
);
let auth_context = resolve_data_backed_auth_context(
&state,
&headers,
&uri("/v1/chat/completions"),
Some("openai:chat"),
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(
auth_context.local_rejection,
Some(GatewayLocalAuthRejection::BalanceDenied {
remaining: Some(0.0),
})
);
assert!(!auth_context.access_allowed);
}
#[tokio::test]
async fn execution_runtime_auth_context_revalidates_cached_wallet_state() {
let api_key = "sk-test-runtime-wallet-cache";
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key(api_key)),
sample_snapshot("key-runtime-wallet-cache", "user-runtime-wallet-cache"),
)]));
let wallet_repository = Arc::new(InMemoryWalletRepository::seed(vec![
StoredWalletSnapshot::new(
"wallet-runtime-cache".to_string(),
Some("user-runtime-wallet-cache".to_string()),
None,
10.0,
0.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
10.0,
0.0,
0.0,
0.0,
100,
)
.expect("wallet should build"),
]));
let data = GatewayDataState::with_auth_and_wallet_for_tests(
auth_repository,
Arc::clone(&wallet_repository),
);
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(data);
let decision = GatewayControlDecision::synthetic(
"/v1/chat/completions",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
);
let mut headers = HeaderMap::new();
headers.insert("x-api-key", api_key.parse().unwrap());
let first = resolve_execution_runtime_auth_context(
&state,
&decision,
&headers,
&uri("/v1/chat/completions"),
"trace-runtime-wallet-cache",
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert!(first.access_allowed);
wallet_repository
.update_auth_user_wallet_snapshot(
"user-runtime-wallet-cache",
0.0,
0.0,
"finite",
"USD",
"active",
10.0,
10.0,
0.0,
0.0,
Some(101),
)
.await
.expect("wallet update should succeed")
.expect("wallet should exist");
let second = resolve_execution_runtime_auth_context(
&state,
&decision,
&headers,
&uri("/v1/chat/completions"),
"trace-runtime-wallet-cache",
)
.await
.expect("resolution should succeed")
.expect("auth context should exist");
assert_eq!(
second.local_rejection,
Some(GatewayLocalAuthRejection::BalanceDenied {
remaining: Some(0.0),
})
);
assert!(!second.access_allowed);
}
#[tokio::test] #[tokio::test]
async fn data_backed_auth_context_allows_provider_id_for_matching_provider_type() { async fn data_backed_auth_context_allows_provider_id_for_matching_provider_type() {
let api_key = "sk-test-provider-id"; let api_key = "sk-test-provider-id";

View File

@@ -137,6 +137,11 @@ const PERMISSION_GROUPS: &[PermissionGroup] = &[
label: "代理节点", label: "代理节点",
assignable: true, assignable: true,
}, },
PermissionGroup {
scope: "routing_profiles",
label: "调度分组",
assignable: true,
},
PermissionGroup { PermissionGroup {
scope: "security", scope: "security",
label: "安全", label: "安全",
@@ -446,6 +451,9 @@ fn permission_key(scope: &str, access: &str) -> &'static str {
("proxy_nodes", "read") => "admin:proxy_nodes:read", ("proxy_nodes", "read") => "admin:proxy_nodes:read",
("proxy_nodes", "write") => "admin:proxy_nodes:write", ("proxy_nodes", "write") => "admin:proxy_nodes:write",
("proxy_nodes", "admin") => "admin:proxy_nodes:admin", ("proxy_nodes", "admin") => "admin:proxy_nodes:admin",
("routing_profiles", "read") => "admin:routing_profiles:read",
("routing_profiles", "write") => "admin:routing_profiles:write",
("routing_profiles", "admin") => "admin:routing_profiles:admin",
("security", "read") => "admin:security:read", ("security", "read") => "admin:security:read",
("security", "write") => "admin:security:write", ("security", "write") => "admin:security:write",
("security", "admin") => "admin:security:admin", ("security", "admin") => "admin:security:admin",

View File

@@ -8,9 +8,10 @@ mod public;
mod route; mod route;
pub(crate) use auth::{ pub(crate) use auth::{
extract_requested_model, request_model_local_rejection, resolve_execution_runtime_auth_context, extract_requested_model, refresh_execution_runtime_auth_context, request_model_local_rejection,
should_buffer_request_for_local_auth, trusted_auth_local_rejection, resolve_execution_runtime_auth_context, should_buffer_request_for_local_auth,
GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayLocalAuthRejection, trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext,
GatewayLocalAuthRejection,
}; };
pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control}; pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control};
pub(crate) use management_token_permissions::{ pub(crate) use management_token_permissions::{

View File

@@ -14,6 +14,8 @@ mod observability_families;
mod operations_families; mod operations_families;
#[path = "admin/provider_ops_routes.rs"] #[path = "admin/provider_ops_routes.rs"]
mod provider_ops_routes; mod provider_ops_routes;
#[path = "admin/routing_families.rs"]
mod routing_families;
#[path = "admin/system_families.rs"] #[path = "admin/system_families.rs"]
mod system_families; mod system_families;
@@ -23,6 +25,7 @@ use model_provider_families::classify_admin_model_provider_family_route;
use observability_families::classify_admin_observability_family_route; use observability_families::classify_admin_observability_family_route;
use operations_families::classify_admin_operations_family_route; use operations_families::classify_admin_operations_family_route;
use provider_ops_routes::classify_admin_provider_ops_routes; use provider_ops_routes::classify_admin_provider_ops_routes;
use routing_families::classify_admin_routing_family_route;
use system_families::classify_admin_system_family_route; use system_families::classify_admin_system_family_route;
pub(super) fn classify_admin_route( pub(super) fn classify_admin_route(
@@ -67,6 +70,10 @@ pub(super) fn classify_admin_route(
classify_admin_system_family_route(method, normalized_path, normalized_path_no_trailing) classify_admin_system_family_route(method, normalized_path, normalized_path_no_trailing)
{ {
Some(route) Some(route)
} else if let Some(route) =
classify_admin_routing_family_route(method, normalized_path_no_trailing)
{
Some(route)
} else if let Some(route) = classify_admin_provider_ops_routes(method, normalized_path) { } else if let Some(route) = classify_admin_provider_ops_routes(method, normalized_path) {
Some(route) Some(route)
} else if let Some(route) = classify_admin_model_provider_family_route(method, normalized_path) } else if let Some(route) = classify_admin_model_provider_family_route(method, normalized_path)

View File

@@ -8,6 +8,56 @@ pub(super) fn classify_admin_operations_family_route(
normalized_path_no_trailing: &str, normalized_path_no_trailing: &str,
) -> Option<ClassifiedRoute> { ) -> Option<ClassifiedRoute> {
if method == http::Method::GET if method == http::Method::GET
&& matches!(
normalized_path,
"/api/admin/referrals" | "/api/admin/referrals/"
)
{
Some(classified(
"admin_proxy",
"referrals_manage",
"list_referrals",
"admin:billing",
false,
))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/admin/referral-rewards" | "/api/admin/referral-rewards/"
)
{
Some(classified(
"admin_proxy",
"referrals_manage",
"list_referral_rewards",
"admin:billing",
false,
))
} else if method == http::Method::POST
&& normalized_path.starts_with("/api/admin/referral-rewards/")
&& normalized_path.ends_with("/retry")
&& normalized_path.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"referrals_manage",
"retry_referral_reward",
"admin:billing",
false,
))
} else if method == http::Method::POST
&& normalized_path.starts_with("/api/admin/referral-rewards/")
&& normalized_path.ends_with("/void")
&& normalized_path.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"referrals_manage",
"void_referral_reward",
"admin:billing",
false,
))
} else if method == http::Method::GET
&& matches!( && matches!(
normalized_path, normalized_path,
"/api/admin/provider-ops/architectures" | "/api/admin/provider-ops/architectures/" "/api/admin/provider-ops/architectures" | "/api/admin/provider-ops/architectures/"

View File

@@ -0,0 +1,74 @@
use axum::http;
use super::{classified, ClassifiedRoute};
pub(super) fn classify_admin_routing_family_route(
method: &http::Method,
normalized_path_no_trailing: &str,
) -> Option<ClassifiedRoute> {
let path = normalized_path_no_trailing;
if method == http::Method::GET && path == "/api/admin/routing/groups" {
Some(routing_route("list_groups"))
} else if method == http::Method::POST && path == "/api/admin/routing/groups" {
Some(routing_route("create_group"))
} else if method == http::Method::GET
&& path.starts_with("/api/admin/routing/groups/")
&& path.ends_with("/versions")
&& path.matches('/').count() == 6
{
Some(routing_route("list_group_versions"))
} else if method == http::Method::POST
&& path.starts_with("/api/admin/routing/groups/")
&& path.ends_with("/publish")
&& path.matches('/').count() == 6
{
Some(routing_route("publish_group"))
} else if method == http::Method::POST
&& path.starts_with("/api/admin/routing/groups/")
&& path.ends_with("/dry-run")
&& path.matches('/').count() == 6
{
Some(routing_route("dry_run_group"))
} else if method == http::Method::GET
&& path.starts_with("/api/admin/routing/groups/")
&& path.matches('/').count() == 5
{
Some(routing_route("get_group"))
} else if method == http::Method::PATCH
&& path.starts_with("/api/admin/routing/groups/")
&& path.matches('/').count() == 5
{
Some(routing_route("update_group"))
} else if method == http::Method::DELETE
&& path.starts_with("/api/admin/routing/groups/")
&& path.matches('/').count() == 5
{
Some(routing_route("delete_group"))
} else if method == http::Method::GET && path == "/api/admin/routing/bindings" {
Some(routing_route("list_bindings"))
} else if method == http::Method::POST && path == "/api/admin/routing/bindings" {
Some(routing_route("create_binding"))
} else if method == http::Method::PATCH
&& path.starts_with("/api/admin/routing/bindings/")
&& path.matches('/').count() == 5
{
Some(routing_route("update_binding"))
} else if method == http::Method::DELETE
&& path.starts_with("/api/admin/routing/bindings/")
&& path.matches('/').count() == 5
{
Some(routing_route("delete_binding"))
} else {
None
}
}
fn routing_route(route_kind: &'static str) -> ClassifiedRoute {
classified(
"admin_proxy",
"routing_profiles_manage",
route_kind,
"admin:routing_profiles",
false,
)
}

View File

@@ -55,7 +55,7 @@ pub(super) fn classify_ai_public_route(
} else if method == http::Method::POST } else if method == http::Method::POST
&& matches!( && matches!(
normalized_path, normalized_path,
"/v1/images/generations" | "/v1/images/edits" | "/v1/images/variations" "/v1/images/generations" | "/v1/images/edits"
) )
{ {
Some(classified( Some(classified(

View File

@@ -279,6 +279,20 @@ pub(super) fn classify_public_support_route(
"user:announcements", "user:announcements",
false, false,
)) ))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/announcements/users/me/required-unread"
| "/api/announcements/users/me/required-unread/"
)
{
Some(classified(
"public_support",
"announcement_user",
"required_unread",
"user:announcements",
false,
))
} else if method == http::Method::POST } else if method == http::Method::POST
&& matches!( && matches!(
normalized_path, normalized_path,
@@ -462,6 +476,7 @@ pub(super) fn classify_public_support_route(
| "/api/users/me/available-models" | "/api/users/me/available-models"
| "/api/users/me/endpoint-status" | "/api/users/me/endpoint-status"
| "/api/users/me/preferences" | "/api/users/me/preferences"
| "/api/users/me/referral"
| "/api/users/me/model-capabilities" | "/api/users/me/model-capabilities"
) )
{ {
@@ -477,6 +492,7 @@ pub(super) fn classify_public_support_route(
"/api/users/me/available-models" => "available_models", "/api/users/me/available-models" => "available_models",
"/api/users/me/endpoint-status" => "endpoint_status", "/api/users/me/endpoint-status" => "endpoint_status",
"/api/users/me/preferences" => "preferences", "/api/users/me/preferences" => "preferences",
"/api/users/me/referral" => "referral",
"/api/users/me/model-capabilities" => "model_capabilities", "/api/users/me/model-capabilities" => "model_capabilities",
_ => "detail", _ => "detail",
}; };

View File

@@ -0,0 +1,92 @@
use http::Uri;
use crate::handlers::shared::local_proxy_route_requires_buffered_body;
use super::{classify_control_route, headers, GatewayPublicRequestContext};
#[test]
fn classifies_admin_routing_group_routes_as_admin_proxy_route() {
let headers = headers(&[]);
let list_uri: Uri = "/api/admin/routing/groups"
.parse()
.expect("uri should parse");
let list = classify_control_route(&http::Method::GET, &list_uri, &headers)
.expect("route should classify");
assert_eq!(list.route_class.as_deref(), Some("admin_proxy"));
assert_eq!(
list.route_family.as_deref(),
Some("routing_profiles_manage")
);
assert_eq!(list.route_kind.as_deref(), Some("list_groups"));
assert_eq!(
list.auth_endpoint_signature.as_deref(),
Some("admin:routing_profiles")
);
let create_uri: Uri = "/api/admin/routing/groups"
.parse()
.expect("uri should parse");
let create = classify_control_route(&http::Method::POST, &create_uri, &headers)
.expect("route should classify");
assert_eq!(
create.route_family.as_deref(),
Some("routing_profiles_manage")
);
assert_eq!(create.route_kind.as_deref(), Some("create_group"));
let update_uri: Uri = "/api/admin/routing/groups/group-1"
.parse()
.expect("uri should parse");
let update = classify_control_route(&http::Method::PATCH, &update_uri, &headers)
.expect("route should classify");
assert_eq!(
update.route_family.as_deref(),
Some("routing_profiles_manage")
);
assert_eq!(update.route_kind.as_deref(), Some("update_group"));
let dry_run_uri: Uri = "/api/admin/routing/groups/group-1/dry-run"
.parse()
.expect("uri should parse");
let dry_run = classify_control_route(&http::Method::POST, &dry_run_uri, &headers)
.expect("route should classify");
assert_eq!(
dry_run.route_family.as_deref(),
Some("routing_profiles_manage")
);
assert_eq!(dry_run.route_kind.as_deref(), Some("dry_run_group"));
}
#[test]
fn admin_routing_write_routes_buffer_request_body() {
let headers = headers(&[]);
let routes = [
(http::Method::POST, "/api/admin/routing/groups"),
(http::Method::PATCH, "/api/admin/routing/groups/group-1"),
(
http::Method::POST,
"/api/admin/routing/groups/group-1/dry-run",
),
(http::Method::POST, "/api/admin/routing/bindings"),
(http::Method::PATCH, "/api/admin/routing/bindings/binding-1"),
];
for (method, path) in routes {
let uri: Uri = path.parse().expect("uri should parse");
let decision =
classify_control_route(&method, &uri, &headers).expect("route should classify");
let context = GatewayPublicRequestContext::from_request_parts(
"trace-routing-write",
&method,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"{method} {path} should buffer request body"
);
}
}

View File

@@ -80,6 +80,28 @@ fn classifies_openai_chat_and_responses_separately_from_embedding() {
assert_ne!(responses.route_kind.as_deref(), Some("embedding")); assert_ne!(responses.route_kind.as_deref(), Some("embedding"));
} }
#[test]
fn classifies_openai_image_generation_and_edit_but_not_variation() {
let headers = headers(&[("authorization", "Bearer sk-test")]);
for path in ["/v1/images/generations", "/v1/images/edits"] {
let uri: Uri = path.parse().expect("uri should parse");
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
.expect("image route should classify");
assert_eq!(decision.route_family.as_deref(), Some("openai"));
assert_eq!(decision.route_kind.as_deref(), Some("image"));
assert_eq!(
decision.auth_endpoint_signature.as_deref(),
Some("openai:image")
);
assert!(decision.is_execution_runtime_candidate());
}
let variation_uri: Uri = "/v1/images/variations".parse().expect("uri should parse");
assert!(classify_control_route(&http::Method::POST, &variation_uri, &headers).is_none());
}
#[test] #[test]
fn classifies_models_list_as_claude_when_headers_match() { fn classifies_models_list_as_claude_when_headers_match() {
let headers = headers(&[ let headers = headers(&[

View File

@@ -86,6 +86,7 @@ mod admin_provider_query;
mod admin_provider_strategy; mod admin_provider_strategy;
mod admin_providers_models; mod admin_providers_models;
mod admin_proxy_nodes; mod admin_proxy_nodes;
mod admin_routing;
mod admin_security; mod admin_security;
mod admin_stats; mod admin_stats;
mod admin_usage; mod admin_usage;

View File

@@ -1,15 +1,15 @@
use super::{ use super::{
AuthApiKeyLookupKey, CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot, AuthApiKeyLookupKey, CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot,
GatewayDataState, ManagementTokenListQuery, ProxyNodeHeartbeatMutation, GatewayDataState, ManagementTokenCounterDelta, ManagementTokenListQuery, ProxyNodeCounterDelta,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeRegistrationMutation, ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation, ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, ProxyNodeTunnelStatusMutation, RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord,
StoredLdapModuleConfig, StoredManagementToken, StoredManagementTokenListPage, StoredAuthApiKeySnapshot, StoredLdapModuleConfig, StoredManagementToken,
StoredManagementTokenWithUser, StoredOAuthProviderConfig, StoredOAuthProviderModuleConfig, StoredManagementTokenListPage, StoredManagementTokenWithUser, StoredOAuthProviderConfig,
StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent, StoredOAuthProviderModuleConfig, StoredProxyFleetMetricsBucket, StoredProxyNode,
StoredProxyNodeMetricsBucket, StoredUserAuthRecord, StoredUserOAuthLinkSummary, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, StoredUserAuthRecord,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredWalletSnapshot, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord, StoredWalletSnapshot, UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord,
}; };
use crate::LocalMutationOutcome; use crate::LocalMutationOutcome;
use aether_data::repository::auth::{ use aether_data::repository::auth::{
@@ -1117,6 +1117,20 @@ impl GatewayDataState {
token_id: &str, token_id: &str,
last_used_ip: Option<&str>, last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> { ) -> Result<Option<StoredManagementToken>, DataLayerError> {
if let Some(repository) = &self.usage_writer {
let enqueued = repository
.enqueue_management_token_counter_delta(ManagementTokenCounterDelta {
token_id: token_id.to_string(),
usage_count_delta: 1,
last_used_at_unix_secs: Some(chrono::Utc::now().timestamp().max(0) as u64),
last_used_ip: last_used_ip.map(ToOwned::to_owned),
})
.await?;
if enqueued {
return Ok(None);
}
}
match &self.management_token_writer { match &self.management_token_writer {
Some(repository) => { Some(repository) => {
repository repository
@@ -1278,6 +1292,21 @@ impl GatewayDataState {
&self, &self,
mutation: &ProxyNodeTrafficMutation, mutation: &ProxyNodeTrafficMutation,
) -> Result<bool, DataLayerError> { ) -> Result<bool, DataLayerError> {
if let Some(repository) = &self.usage_writer {
let enqueued = repository
.enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta {
node_id: mutation.node_id.clone(),
total_requests_delta: mutation.total_requests_delta,
failed_requests_delta: mutation.failed_requests_delta,
dns_failures_delta: mutation.dns_failures_delta,
stream_errors_delta: mutation.stream_errors_delta,
})
.await?;
if enqueued {
return Ok(true);
}
}
match &self.proxy_node_writer { match &self.proxy_node_writer {
Some(repository) => repository.record_traffic(mutation).await, Some(repository) => repository.record_traffic(mutation).await,
None => Ok(false), None => Ok(false),
@@ -1925,14 +1954,38 @@ fn resolve_effective_list_policy(
&aether_data::repository::users::StoredUserGroup, &aether_data::repository::users::StoredUserGroup,
) -> (&str, Option<Vec<String>>), ) -> (&str, Option<Vec<String>>),
) -> Option<Vec<String>> { ) -> Option<Vec<String>> {
let group_policy = groups.iter().fold(None, |effective, group| { let group_policy = union_group_list_policies(groups, group_field);
let (mode, values) = group_field(group);
intersect_list_policies(effective, list_restriction_from_mode(mode, values))
});
let user_policy = list_restriction_from_mode(user_mode, user_values); let user_policy = list_restriction_from_mode(user_mode, user_values);
intersect_list_policies(group_policy, user_policy) intersect_list_policies(group_policy, user_policy)
} }
fn union_group_list_policies(
groups: &[aether_data::repository::users::StoredUserGroup],
group_field: impl Fn(
&aether_data::repository::users::StoredUserGroup,
) -> (&str, Option<Vec<String>>),
) -> Option<Vec<String>> {
let mut saw_restrictive_group = false;
let mut values = std::collections::BTreeSet::new();
for group in groups {
let (mode, group_values) = group_field(group);
match mode {
"unrestricted" => return None,
"specific" => {
saw_restrictive_group = true;
values.extend(group_values.unwrap_or_default());
}
"deny_all" => {
saw_restrictive_group = true;
}
_ => {}
}
}
saw_restrictive_group.then(|| values.into_iter().collect())
}
fn list_restriction_from_mode(mode: &str, values: Option<Vec<String>>) -> Option<Vec<String>> { fn list_restriction_from_mode(mode: &str, values: Option<Vec<String>>) -> Option<Vec<String>> {
match mode { match mode {
"specific" => Some(values.unwrap_or_default()), "specific" => Some(values.unwrap_or_default()),
@@ -2148,7 +2201,7 @@ mod tests {
} }
#[test] #[test]
fn list_policy_intersects_group_and_user_restrictions() { fn list_policy_intersects_unrestricted_group_union_with_user_restriction() {
let groups = vec![ let groups = vec![
sample_group("default", 0, None, "unrestricted", None, "system"), sample_group("default", 0, None, "unrestricted", None, "system"),
sample_group( sample_group(
@@ -2168,11 +2221,14 @@ mod tests {
|group| (&group.allowed_models_mode, group.allowed_models.clone()), |group| (&group.allowed_models_mode, group.allowed_models.clone()),
); );
assert_eq!(policy, Some(vec!["gpt-4.1".to_string()])); assert_eq!(
policy,
Some(vec!["gpt-4.1".to_string(), "gemini-2.5-pro".to_string()])
);
} }
#[test] #[test]
fn list_policy_intersects_multiple_group_restrictions() { fn list_policy_unions_multiple_group_restrictions_legacy_case() {
let groups = vec![ let groups = vec![
sample_group( sample_group(
"team-a", "team-a",
@@ -2196,7 +2252,91 @@ mod tests {
(&group.allowed_models_mode, group.allowed_models.clone()) (&group.allowed_models_mode, group.allowed_models.clone())
}); });
assert_eq!(policy, Some(vec!["gpt-4.1".to_string()])); assert_eq!(
policy,
Some(vec![
"gemini-2.5-pro".to_string(),
"gpt-4.1".to_string(),
"gpt-5".to_string()
])
);
}
#[test]
fn list_policy_unions_multiple_group_restrictions() {
let groups = vec![
sample_group(
"team-a",
10,
Some(vec!["gpt-5", "gpt-4.1"]),
"specific",
None,
"system",
),
sample_group(
"team-b",
20,
Some(vec!["gpt-4.1", "gemini-2.5-pro"]),
"specific",
None,
"system",
),
];
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
(&group.allowed_models_mode, group.allowed_models.clone())
});
assert_eq!(
policy,
Some(vec![
"gemini-2.5-pro".to_string(),
"gpt-4.1".to_string(),
"gpt-5".to_string()
])
);
}
#[test]
fn unrestricted_group_makes_group_policy_unrestricted() {
let groups = vec![
sample_group(
"restricted",
10,
Some(vec!["gpt-5"]),
"specific",
None,
"system",
),
sample_group("unrestricted", 20, None, "unrestricted", None, "system"),
];
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
(&group.allowed_models_mode, group.allowed_models.clone())
});
assert_eq!(policy, None);
}
#[test]
fn deny_all_group_does_not_remove_other_group_grants() {
let groups = vec![
sample_group("deny", 10, None, "deny_all", None, "system"),
sample_group(
"restricted",
20,
Some(vec!["gpt-5"]),
"specific",
None,
"system",
),
];
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
(&group.allowed_models_mode, group.allowed_models.clone())
});
assert_eq!(policy, Some(vec!["gpt-5".to_string()]));
} }
#[test] #[test]

View File

@@ -0,0 +1,401 @@
use std::collections::HashSet;
use std::future::Future;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use aether_cache::ExpiringMap;
use aether_data::DataLayerError;
use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use async_trait::async_trait;
use tokio::sync::Notify;
const CANDIDATE_SELECTION_CACHE_TTL: Duration = Duration::from_secs(5);
const CANDIDATE_SELECTION_CACHE_MAX_ENTRIES: usize = 4096;
pub(super) struct CachedMinimalCandidateSelectionReadRepository {
inner: Arc<dyn MinimalCandidateSelectionReadRepository>,
entries: ExpiringMap<CandidateSelectionCacheKey, Vec<StoredMinimalCandidateSelectionRow>>,
inflight: Mutex<HashSet<CandidateSelectionCacheKey>>,
inflight_notify: Notify,
epoch: AtomicU64,
}
impl CachedMinimalCandidateSelectionReadRepository {
pub(super) fn new(inner: Arc<dyn MinimalCandidateSelectionReadRepository>) -> Self {
Self {
inner,
entries: ExpiringMap::new(),
inflight: Mutex::new(HashSet::new()),
inflight_notify: Notify::new(),
epoch: AtomicU64::new(0),
}
}
async fn get_or_load<F, Fut>(
&self,
key: CandidateSelectionCacheKey,
load: F,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>>,
{
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL) {
return Ok(rows);
}
loop {
let notified = self.inflight_notify.notified();
match self.register_inflight(&key) {
InflightRegistration::Bypass => return load().await,
InflightRegistration::Follower => {
notified.await;
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
{
return Ok(rows);
}
continue;
}
InflightRegistration::Leader => {}
}
let load_epoch = self.epoch.load(Ordering::Acquire);
let result = load().await;
if let Ok(rows) = &result {
if load_epoch == self.epoch.load(Ordering::Acquire) {
self.entries.insert(
key.clone(),
rows.clone(),
CANDIDATE_SELECTION_CACHE_TTL,
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
);
}
}
self.finish_inflight(&key);
return result;
}
}
fn register_inflight(&self, key: &CandidateSelectionCacheKey) -> InflightRegistration {
match self.inflight.lock() {
Ok(mut inflight) => {
if inflight.insert(key.clone()) {
InflightRegistration::Leader
} else {
InflightRegistration::Follower
}
}
Err(_) => InflightRegistration::Bypass,
}
}
fn finish_inflight(&self, key: &CandidateSelectionCacheKey) {
if let Ok(mut inflight) = self.inflight.lock() {
inflight.remove(key);
}
self.inflight_notify.notify_waiters();
}
fn clear(&self) {
self.epoch.fetch_add(1, Ordering::AcqRel);
self.entries.clear();
}
}
enum InflightRegistration {
Leader,
Follower,
Bypass,
}
#[async_trait]
impl MinimalCandidateSelectionReadRepository for CachedMinimalCandidateSelectionReadRepository {
fn clear_local_cache(&self) {
self.clear();
self.inner.clear_local_cache();
}
async fn list_for_exact_api_format(
&self,
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::ApiFormat {
api_format: normalize_api_format_key(api_format),
};
self.get_or_load(key, || self.inner.list_for_exact_api_format(api_format))
.await
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::ApiFormatAndGlobalModel {
api_format: normalize_api_format_key(api_format),
global_model_name: global_model_name.to_string(),
};
self.get_or_load(key, || {
self.inner
.list_for_exact_api_format_and_global_model(api_format, global_model_name)
})
.await
}
async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::ApiFormatAndRequestedModel {
api_format: normalize_api_format_key(api_format),
requested_model_name: requested_model_name.to_string(),
};
self.get_or_load(key, || {
self.inner
.list_for_exact_api_format_and_requested_model(api_format, requested_model_name)
})
.await
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::RequestedModelPage {
api_format: normalize_api_format_key(&query.api_format),
requested_model_name: query.requested_model_name.clone(),
offset: query.offset,
limit: query.limit,
};
self.get_or_load(key, || {
self.inner
.list_for_exact_api_format_and_requested_model_page(query)
})
.await
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::PoolKeyRowsForGroup {
api_format: normalize_api_format_key(&query.api_format),
provider_id: query.provider_id.clone(),
endpoint_id: query.endpoint_id.clone(),
model_id: query.model_id.clone(),
selected_provider_model_name: query.selected_provider_model_name.clone(),
order: CandidateSelectionPoolOrderKey::from(&query.order),
offset: query.offset,
limit: query.limit,
};
self.get_or_load(key, || self.inner.list_pool_key_rows_for_group(query))
.await
}
async fn list_pool_key_rows_for_group_key_ids(
&self,
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let key = CandidateSelectionCacheKey::PoolKeyRowsForGroupKeyIds {
api_format: normalize_api_format_key(&query.api_format),
provider_id: query.provider_id.clone(),
endpoint_id: query.endpoint_id.clone(),
model_id: query.model_id.clone(),
selected_provider_model_name: query.selected_provider_model_name.clone(),
key_ids: query.key_ids.clone(),
};
self.get_or_load(key, || {
self.inner.list_pool_key_rows_for_group_key_ids(query)
})
.await
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum CandidateSelectionCacheKey {
ApiFormat {
api_format: String,
},
ApiFormatAndGlobalModel {
api_format: String,
global_model_name: String,
},
ApiFormatAndRequestedModel {
api_format: String,
requested_model_name: String,
},
RequestedModelPage {
api_format: String,
requested_model_name: String,
offset: u32,
limit: u32,
},
PoolKeyRowsForGroup {
api_format: String,
provider_id: String,
endpoint_id: String,
model_id: String,
selected_provider_model_name: String,
order: CandidateSelectionPoolOrderKey,
offset: u32,
limit: u32,
},
PoolKeyRowsForGroupKeyIds {
api_format: String,
provider_id: String,
endpoint_id: String,
model_id: String,
selected_provider_model_name: String,
key_ids: Vec<String>,
},
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
enum CandidateSelectionPoolOrderKey {
InternalPriority,
Lru,
CacheAffinity,
SingleAccount,
LoadBalance { seed: String },
}
impl From<&StoredPoolKeyCandidateOrder> for CandidateSelectionPoolOrderKey {
fn from(order: &StoredPoolKeyCandidateOrder) -> Self {
match order {
StoredPoolKeyCandidateOrder::InternalPriority => Self::InternalPriority,
StoredPoolKeyCandidateOrder::Lru => Self::Lru,
StoredPoolKeyCandidateOrder::CacheAffinity => Self::CacheAffinity,
StoredPoolKeyCandidateOrder::SingleAccount => Self::SingleAccount,
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
Self::LoadBalance { seed: seed.clone() }
}
}
}
}
fn normalize_api_format_key(api_format: &str) -> String {
crate::ai_serving::normalize_api_format_alias(api_format.trim())
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::AtomicUsize;
struct StubCandidateSelectionRepository {
calls: AtomicUsize,
delay: Duration,
}
impl StubCandidateSelectionRepository {
fn new(delay: Duration) -> Self {
Self {
calls: AtomicUsize::new(0),
delay,
}
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
async fn load(&self) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.calls.fetch_add(1, Ordering::SeqCst);
if !self.delay.is_zero() {
tokio::time::sleep(self.delay).await;
}
Ok(Vec::new())
}
}
#[async_trait]
impl MinimalCandidateSelectionReadRepository for StubCandidateSelectionRepository {
async fn list_for_exact_api_format(
&self,
_api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
async fn list_for_exact_api_format_and_global_model(
&self,
_api_format: &str,
_global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
async fn list_for_exact_api_format_and_requested_model(
&self,
_api_format: &str,
_requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
_query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
async fn list_pool_key_rows_for_group(
&self,
_query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
async fn list_pool_key_rows_for_group_key_ids(
&self,
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.load().await
}
}
#[tokio::test]
async fn candidate_selection_cache_coalesces_concurrent_loads() {
let inner = Arc::new(StubCandidateSelectionRepository::new(
Duration::from_millis(25),
));
let cache = Arc::new(CachedMinimalCandidateSelectionReadRepository::new(
inner.clone(),
));
let mut tasks = Vec::new();
for _ in 0..16 {
let cache = cache.clone();
tasks.push(tokio::spawn(async move {
cache.list_for_exact_api_format("openai").await.unwrap();
}));
}
for task in tasks {
task.await.unwrap();
}
assert_eq!(inner.calls(), 1);
}
#[tokio::test]
async fn candidate_selection_cache_clear_invalidates_entries() {
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner.clone());
cache.list_for_exact_api_format("openai").await.unwrap();
cache.list_for_exact_api_format("openai").await.unwrap();
assert_eq!(inner.calls(), 1);
cache.clear_local_cache();
cache.list_for_exact_api_format("openai").await.unwrap();
assert_eq!(inner.calls(), 2);
}
}

View File

@@ -1,10 +1,10 @@
use super::{ use super::{
DataLayerError, GatewayDataState, GeminiFileMappingListQuery, GeminiFileMappingStats, ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
ProviderCatalogKeyListQuery, PublicHealthStatusCount, PublicHealthTimelineBucket, GeminiFileMappingStats, ProviderCatalogKeyListQuery, PublicHealthStatusCount,
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint, PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage,
StoredProviderCatalogKey, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
StoredProviderCatalogProvider, StoredRequestCandidate, UpsertGeminiFileMappingRecord, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
UpsertRequestCandidateRecord, UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
}; };
impl GatewayDataState { impl GatewayDataState {
@@ -121,6 +121,18 @@ impl GatewayDataState {
&self, &self,
api_key_id: &str, api_key_id: &str,
) -> Result<bool, DataLayerError> { ) -> Result<bool, DataLayerError> {
if let Some(repository) = &self.usage_writer {
let enqueued = repository
.enqueue_api_key_last_used_delta(ApiKeyLastUsedDelta {
api_key_id: api_key_id.to_string(),
last_used_at_unix_secs: chrono::Utc::now().timestamp().max(0) as u64,
})
.await?;
if enqueued {
return Ok(true);
}
}
match &self.auth_api_key_writer { match &self.auth_api_key_writer {
Some(repository) => repository.touch_last_used_at(api_key_id).await, Some(repository) => repository.touch_last_used_at(api_key_id).await,
None => Ok(false), None => Ok(false),

View File

@@ -1,4 +1,5 @@
use aether_data::{DataBackends, DataLayerError, DatabaseDriver}; use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
use aether_runtime_state::RuntimeQueueStore; use aether_runtime_state::RuntimeQueueStore;
use std::sync::Arc; use std::sync::Arc;
@@ -49,6 +50,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -82,7 +85,17 @@ impl GatewayDataState {
let gemini_file_mapping_reader = backends.read().gemini_file_mappings(); let gemini_file_mapping_reader = backends.read().gemini_file_mappings();
let global_model_reader = backends.read().global_models(); let global_model_reader = backends.read().global_models();
let global_model_writer = backends.write().global_models(); let global_model_writer = backends.write().global_models();
let minimal_candidate_selection_reader = backends.read().minimal_candidate_selection(); let minimal_candidate_selection_reader =
backends
.read()
.minimal_candidate_selection()
.map(|repository| {
Arc::new(
super::candidate_cache::CachedMinimalCandidateSelectionReadRepository::new(
repository,
),
) as Arc<dyn MinimalCandidateSelectionReadRepository>
});
let request_candidate_reader = backends.read().request_candidates(); let request_candidate_reader = backends.read().request_candidates();
let request_candidate_writer = backends.write().request_candidates(); let request_candidate_writer = backends.write().request_candidates();
let gemini_file_mapping_writer = backends.write().gemini_file_mappings(); let gemini_file_mapping_writer = backends.write().gemini_file_mappings();
@@ -92,6 +105,8 @@ impl GatewayDataState {
let pool_score_writer = backends.write().pool_scores(); let pool_score_writer = backends.write().pool_scores();
let provider_quota_reader = backends.read().provider_quotas(); let provider_quota_reader = backends.read().provider_quotas();
let provider_quota_writer = backends.write().provider_quotas(); let provider_quota_writer = backends.write().provider_quotas();
let routing_group_reader = backends.read().routing_groups();
let routing_group_writer = backends.write().routing_groups();
let usage_reader = backends.read().usage(); let usage_reader = backends.read().usage();
let usage_writer = backends.write().usage(); let usage_writer = backends.write().usage();
let user_reader = backends.read().users(); let user_reader = backends.read().users();
@@ -133,6 +148,8 @@ impl GatewayDataState {
pool_score_writer, pool_score_writer,
provider_quota_reader, provider_quota_reader,
provider_quota_writer, provider_quota_writer,
routing_group_reader,
routing_group_writer,
usage_reader, usage_reader,
usage_writer, usage_writer,
user_reader, user_reader,
@@ -253,6 +270,12 @@ impl GatewayDataState {
self.minimal_candidate_selection_reader.is_some() self.minimal_candidate_selection_reader.is_some()
} }
pub(crate) fn clear_minimal_candidate_selection_cache(&self) {
if let Some(repository) = &self.minimal_candidate_selection_reader {
repository.clear_local_cache();
}
}
pub(crate) fn has_request_candidate_reader(&self) -> bool { pub(crate) fn has_request_candidate_reader(&self) -> bool {
self.request_candidate_reader.is_some() self.request_candidate_reader.is_some()
} }
@@ -261,6 +284,14 @@ impl GatewayDataState {
self.request_candidate_writer.is_some() self.request_candidate_writer.is_some()
} }
pub(crate) fn has_routing_group_reader(&self) -> bool {
self.routing_group_reader.is_some()
}
pub(crate) fn has_routing_group_writer(&self) -> bool {
self.routing_group_writer.is_some()
}
pub(crate) fn has_provider_catalog_reader(&self) -> bool { pub(crate) fn has_provider_catalog_reader(&self) -> bool {
self.provider_catalog_reader.is_some() self.provider_catalog_reader.is_some()
} }
@@ -315,6 +346,10 @@ impl GatewayDataState {
self.usage_writer.is_some() self.usage_writer.is_some()
} }
pub(crate) fn has_usage_counter_flush_backend(&self) -> bool {
self.has_usage_writer() && self.database_driver() == Some(DatabaseDriver::Postgres)
}
pub(crate) fn has_usage_worker_queue(&self) -> bool { pub(crate) fn has_usage_worker_queue(&self) -> bool {
self.usage_worker_queue.is_some() self.usage_worker_queue.is_some()
} }

View File

@@ -12,7 +12,9 @@ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput}; use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord}; use aether_data_contracts::repository::usage::{
ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord,
};
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey}; use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey};
use aether_runtime_state::RuntimeQueueStore; use aether_runtime_state::RuntimeQueueStore;
use aether_usage_runtime::{ use aether_usage_runtime::{
@@ -284,6 +286,21 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState {
failed_delta: i64, failed_delta: i64,
latency_ms: Option<i64>, latency_ms: Option<i64>,
) -> Result<(), DataLayerError> { ) -> Result<(), DataLayerError> {
if let Some(repository) = &self.usage_writer {
let enqueued = repository
.enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta {
node_id: node_id.to_string(),
total_requests_delta: total_delta,
failed_requests_delta: failed_delta,
dns_failures_delta: 0,
stream_errors_delta: 0,
})
.await?;
if enqueued {
return Ok(());
}
}
match &self.proxy_node_writer { match &self.proxy_node_writer {
Some(repository) => { Some(repository) => {
repository repository

View File

@@ -125,12 +125,16 @@ use aether_data_contracts::repository::provider_catalog::{
use aether_data_contracts::repository::quota::{ use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot, ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
}; };
use aether_data_contracts::repository::routing_profiles::{
RoutingGroupReadRepository, RoutingGroupWriteRepository,
};
use aether_data_contracts::repository::settlement::{ use aether_data_contracts::repository::settlement::{
SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput, SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
}; };
use aether_data_contracts::repository::usage::{ use aether_data_contracts::repository::usage::{
PendingUsageCleanupSummary, StoredProviderUsageSummary, StoredRequestUsageAudit, ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary,
UpsertUsageRecord, UsageReadRepository, UsageWriteRepository, ProxyNodeCounterDelta, StoredProviderUsageSummary, StoredRequestUsageAudit, UpsertUsageRecord,
UsageReadRepository, UsageWriteRepository,
}; };
use aether_data_contracts::repository::video_tasks::{ use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount, StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
@@ -138,6 +142,12 @@ use aether_data_contracts::repository::video_tasks::{
}; };
use aether_runtime_state::RuntimeQueueStore; use aether_runtime_state::RuntimeQueueStore;
pub(crate) use self::referrals::{
ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery,
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
ReferralRewardRecord, ReferralUserDashboard,
};
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub(crate) struct GatewayDataState { pub(crate) struct GatewayDataState {
config: GatewayDataConfig, config: GatewayDataConfig,
@@ -170,6 +180,8 @@ pub(crate) struct GatewayDataState {
pool_score_writer: Option<Arc<dyn PoolMemberScoreWriteRepository>>, pool_score_writer: Option<Arc<dyn PoolMemberScoreWriteRepository>>,
provider_quota_reader: Option<Arc<dyn ProviderQuotaReadRepository>>, provider_quota_reader: Option<Arc<dyn ProviderQuotaReadRepository>>,
provider_quota_writer: Option<Arc<dyn ProviderQuotaWriteRepository>>, provider_quota_writer: Option<Arc<dyn ProviderQuotaWriteRepository>>,
routing_group_reader: Option<Arc<dyn RoutingGroupReadRepository>>,
routing_group_writer: Option<Arc<dyn RoutingGroupWriteRepository>>,
usage_reader: Option<Arc<dyn UsageReadRepository>>, usage_reader: Option<Arc<dyn UsageReadRepository>>,
usage_writer: Option<Arc<dyn UsageWriteRepository>>, usage_writer: Option<Arc<dyn UsageWriteRepository>>,
user_reader: Option<Arc<dyn UserReadRepository>>, user_reader: Option<Arc<dyn UserReadRepository>>,
@@ -279,6 +291,14 @@ impl fmt::Debug for GatewayDataState {
"has_provider_quota_writer", "has_provider_quota_writer",
&self.provider_quota_writer.is_some(), &self.provider_quota_writer.is_some(),
) )
.field(
"has_routing_group_reader",
&self.routing_group_reader.is_some(),
)
.field(
"has_routing_group_writer",
&self.routing_group_writer.is_some(),
)
.field("has_usage_reader", &self.usage_reader.is_some()) .field("has_usage_reader", &self.usage_reader.is_some())
.field("has_usage_writer", &self.usage_writer.is_some()) .field("has_usage_writer", &self.usage_writer.is_some())
.field("has_user_preferences", &self.user_preferences.is_some()) .field("has_user_preferences", &self.user_preferences.is_some())
@@ -297,11 +317,14 @@ impl fmt::Debug for GatewayDataState {
} }
mod auth; mod auth;
mod candidate_cache;
mod catalog; mod catalog;
mod core; mod core;
mod integrations; mod integrations;
mod models; mod models;
mod pool_scores; mod pool_scores;
mod referrals;
mod routing_profiles;
mod runtime; mod runtime;
#[cfg(test)] #[cfg(test)]
mod testing; mod testing;

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,131 @@
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, CreateRoutingGroupVersionRecord,
RoutingGroupBindingQuery, RoutingGroupLookupKey, RoutingGroupReadRepository,
StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion,
UpdateRoutingGroupBindingRecord, UpdateRoutingGroupRecord,
};
use std::sync::Arc;
use super::{DataLayerError, GatewayDataState};
impl GatewayDataState {
pub(crate) fn routing_group_read_repository(
&self,
) -> Option<Arc<dyn RoutingGroupReadRepository>> {
self.routing_group_reader.clone()
}
pub(crate) async fn list_routing_groups(
&self,
) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
match &self.routing_group_reader {
Some(repository) => repository.list_routing_groups().await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn find_routing_group(
&self,
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
match &self.routing_group_reader {
Some(repository) => repository.find_routing_group(lookup).await,
None => Ok(None),
}
}
pub(crate) async fn list_routing_group_bindings(
&self,
query: &RoutingGroupBindingQuery,
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
match &self.routing_group_reader {
Some(repository) => repository.list_routing_group_bindings(query).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_routing_group_versions(
&self,
group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
match &self.routing_group_reader {
Some(repository) => repository.list_routing_group_versions(group_id).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn create_routing_group(
&self,
record: CreateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository.create_routing_group(record).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn update_routing_group(
&self,
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository.update_routing_group(id, patch).await,
None => Ok(None),
}
}
pub(crate) async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository.delete_routing_group(id).await,
None => Ok(false),
}
}
pub(crate) async fn create_routing_group_binding(
&self,
record: CreateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository
.create_routing_group_binding(record)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn update_routing_group_binding(
&self,
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository.update_routing_group_binding(id, patch).await,
None => Ok(None),
}
}
pub(crate) async fn delete_routing_group_binding(
&self,
id: &str,
) -> Result<bool, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository.delete_routing_group_binding(id).await,
None => Ok(false),
}
}
pub(crate) async fn create_routing_group_version(
&self,
record: CreateRoutingGroupVersionRecord,
) -> Result<Option<StoredRoutingGroupVersion>, DataLayerError> {
match &self.routing_group_writer {
Some(repository) => repository
.create_routing_group_version(record)
.await
.map(Some),
None => Ok(None),
}
}
}

View File

@@ -38,7 +38,7 @@ use aether_data_contracts::repository::usage::{
PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest, PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest,
StoredProviderApiKeyWindowUsageSummary, StoredUsageDailySummary, UsageAuditListQuery, StoredProviderApiKeyWindowUsageSummary, StoredUsageDailySummary, UsageAuditListQuery,
UsageCleanupExecutionMode, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow, UsageCleanupExecutionMode, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow,
UsageDailyHeatmapQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageDailyHeatmapQuery,
}; };
use aether_runtime_state::RuntimeQueueStore; use aether_runtime_state::RuntimeQueueStore;
use aether_video_tasks_core::read_data_backed_video_task_response; use aether_video_tasks_core::read_data_backed_video_task_response;
@@ -262,6 +262,22 @@ impl GatewayDataState {
} }
} }
pub(crate) async fn list_required_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredAnnouncement>, DataLayerError> {
match &self.announcement_reader {
Some(repository) => {
repository
.list_required_unread_active_announcements(user_id, now_unix_secs, limit)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn create_announcement( pub(crate) async fn create_announcement(
&self, &self,
record: CreateAnnouncementRecord, record: CreateAnnouncementRecord,
@@ -954,6 +970,31 @@ impl GatewayDataState {
} }
} }
pub(crate) async fn flush_usage_counter_deltas(
&self,
batch_size: usize,
) -> Result<UsageCounterFlushSummary, DataLayerError> {
match &self.usage_writer {
Some(repository) => repository.flush_usage_counter_deltas(batch_size).await,
None => Ok(UsageCounterFlushSummary::default()),
}
}
pub(crate) async fn cleanup_processed_usage_counter_deltas(
&self,
cutoff_unix_secs: u64,
batch_size: usize,
) -> Result<usize, DataLayerError> {
match &self.usage_writer {
Some(repository) => {
repository
.cleanup_processed_usage_counter_deltas(cutoff_unix_secs, batch_size)
.await
}
None => Ok(0),
}
}
pub(crate) async fn cleanup_stale_pending_requests( pub(crate) async fn cleanup_stale_pending_requests(
&self, &self,
cutoff_unix_secs: u64, cutoff_unix_secs: u64,
@@ -1119,6 +1160,15 @@ impl GatewayDataState {
} }
} }
pub(crate) async fn read_usage_counter_health(
&self,
) -> Result<UsageCounterHealthSnapshot, DataLayerError> {
match &self.usage_reader {
Some(repository) => repository.read_usage_counter_health().await,
None => Ok(UsageCounterHealthSnapshot::default()),
}
}
pub(crate) async fn summarize_usage_totals_by_user_ids( pub(crate) async fn summarize_usage_totals_by_user_ids(
&self, &self,
user_ids: &[String], user_ids: &[String],

View File

@@ -38,6 +38,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -92,6 +94,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,

View File

@@ -73,6 +73,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -126,6 +128,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -175,6 +179,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -312,6 +318,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -389,6 +397,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -447,6 +457,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -514,6 +526,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -563,6 +577,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -613,6 +629,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -674,6 +692,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -737,6 +757,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -784,6 +806,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(repository), usage_reader: Some(repository),
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -846,6 +870,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: Some(repository), user_reader: Some(repository),
@@ -901,6 +927,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: Some(user_repository), user_reader: Some(user_repository),
@@ -961,6 +989,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: Some(user_repository), user_reader: Some(user_repository),
@@ -1022,6 +1052,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: Some(user_repository), user_reader: Some(user_repository),
@@ -1082,6 +1114,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1131,6 +1165,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1180,6 +1216,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1241,6 +1279,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1307,6 +1347,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1356,6 +1398,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1410,6 +1454,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1481,6 +1527,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1547,6 +1595,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1597,6 +1647,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1647,6 +1699,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1699,6 +1753,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_repository), usage_reader: Some(usage_repository),
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1749,6 +1805,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1799,6 +1857,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1849,6 +1909,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1907,6 +1969,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -1966,6 +2030,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2028,6 +2094,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2096,6 +2164,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2165,6 +2235,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2238,6 +2310,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -2318,6 +2392,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -2380,6 +2456,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2433,6 +2511,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -2482,6 +2562,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2537,6 +2619,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -2596,6 +2680,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -2656,6 +2742,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: Some(usage_reader), usage_reader: Some(usage_reader),
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
@@ -2709,6 +2797,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: Some(provider_quota_reader), provider_quota_reader: Some(provider_quota_reader),
provider_quota_writer: Some(provider_quota_writer), provider_quota_writer: Some(provider_quota_writer),
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,

View File

@@ -42,6 +42,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -98,6 +100,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -151,6 +155,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -208,6 +214,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -269,6 +277,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
@@ -339,6 +349,8 @@ impl GatewayDataState {
pool_score_writer: None, pool_score_writer: None,
provider_quota_reader: None, provider_quota_reader: None,
provider_quota_writer: None, provider_quota_writer: None,
routing_group_reader: None,
routing_group_writer: None,
usage_reader: None, usage_reader: None,
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,

View File

@@ -16,6 +16,7 @@ use aether_pool_core::{
PoolMemberSignals, PoolRuntimeState, PoolSchedulingConfig, PoolSchedulingPreset, PoolMemberSignals, PoolRuntimeState, PoolSchedulingConfig, PoolSchedulingPreset,
}; };
use aether_provider_pool::ProviderPoolService; use aether_provider_pool::ProviderPoolService;
use aether_routing_core::{RankingOverlay, ResolvedRoutingPolicy};
use tracing::warn; use tracing::warn;
use crate::ai_serving::{ use crate::ai_serving::{
@@ -40,6 +41,7 @@ use crate::orchestration::LocalExecutionCandidateMetadata;
static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0); static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
const POOL_ACTIVE_PROBE_SEALED_SKIP_REASON: &str = "pool_active_probe_sealed"; const POOL_ACTIVE_PROBE_SEALED_SKIP_REASON: &str = "pool_active_probe_sealed";
const ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON: &str = "routing_profile_disallowed_key";
type PoolCatalogKeyContext = PoolMemberSignals; type PoolCatalogKeyContext = PoolMemberSignals;
@@ -187,6 +189,7 @@ pub(crate) struct PoolKeyCursor<'a> {
sticky_session_token: Option<String>, sticky_session_token: Option<String>,
requested_model: Option<String>, requested_model: Option<String>,
request_auth_channel: Option<String>, request_auth_channel: Option<String>,
routing_overlay: Option<RankingOverlay>,
runtime_miss_trace_id: Option<String>, runtime_miss_trace_id: Option<String>,
record_runtime_miss_diagnostic: bool, record_runtime_miss_diagnostic: bool,
pool_key_order: StoredPoolKeyCandidateOrder, pool_key_order: StoredPoolKeyCandidateOrder,
@@ -216,7 +219,26 @@ impl<'a> PoolKeyCursor<'a> {
requested_model: Option<&str>, requested_model: Option<&str>,
request_auth_channel: Option<&str>, request_auth_channel: Option<&str>,
) -> Self { ) -> Self {
let pool_key_order = pool_key_candidate_order_for_group(&group); Self::new_with_routing_policy(
state,
group,
sticky_session_token,
requested_model,
request_auth_channel,
None,
)
}
pub(crate) fn new_with_routing_policy(
state: PlannerAppState<'a>,
group: EligibleLocalExecutionCandidate,
sticky_session_token: Option<&str>,
requested_model: Option<&str>,
request_auth_channel: Option<&str>,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> Self {
let pool_key_order = pool_key_candidate_order_for_group(&group, routing_policy);
let routing_overlay = routing_policy.map(|policy| policy.ranking_overlay.clone());
let pool_config = pool_config_for_candidate(&group); let pool_config = pool_config_for_candidate(&group);
let score_top_n = pool_config let score_top_n = pool_config
.as_ref() .as_ref()
@@ -236,6 +258,7 @@ impl<'a> PoolKeyCursor<'a> {
sticky_session_token: sticky_session_token.map(str::to_string), sticky_session_token: sticky_session_token.map(str::to_string),
requested_model: requested_model.map(str::to_string), requested_model: requested_model.map(str::to_string),
request_auth_channel: request_auth_channel.map(str::to_string), request_auth_channel: request_auth_channel.map(str::to_string),
routing_overlay,
runtime_miss_trace_id: None, runtime_miss_trace_id: None,
record_runtime_miss_diagnostic: false, record_runtime_miss_diagnostic: false,
pool_key_order, pool_key_order,
@@ -551,6 +574,9 @@ impl<'a> PoolKeyCursor<'a> {
async fn next_queued_candidate(&mut self) -> Option<EligibleLocalExecutionCandidate> { async fn next_queued_candidate(&mut self) -> Option<EligibleLocalExecutionCandidate> {
while let Some(candidate) = self.queued_candidates.pop_front() { while let Some(candidate) = self.queued_candidates.pop_front() {
let mut candidate = candidate; let mut candidate = candidate;
if self.skip_candidate_if_routing_profile_disallowed(&candidate) {
continue;
}
if self.skip_candidate_if_runtime_cooldown(&candidate).await { if self.skip_candidate_if_runtime_cooldown(&candidate).await {
continue; continue;
} }
@@ -562,6 +588,28 @@ impl<'a> PoolKeyCursor<'a> {
None None
} }
fn skip_candidate_if_routing_profile_disallowed(
&mut self,
candidate: &EligibleLocalExecutionCandidate,
) -> bool {
let Some(overlay) = self.routing_overlay.as_ref() else {
return false;
};
if overlay.key_allowed(candidate.candidate.key_id.as_str()) {
return false;
}
self.record_skip_reason(ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON);
self.skipped_candidates
.push(SkippedLocalExecutionCandidate {
candidate: candidate.candidate.clone(),
skip_reason: ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON,
transport: Some(candidate.transport.clone()),
ranking: candidate.ranking.clone(),
extra_data: None,
});
true
}
async fn skip_candidate_if_runtime_cooldown( async fn skip_candidate_if_runtime_cooldown(
&mut self, &mut self,
candidate: &EligibleLocalExecutionCandidate, candidate: &EligibleLocalExecutionCandidate,
@@ -958,19 +1006,38 @@ fn should_trigger_active_probe_burst_for_request(
fn pool_key_candidate_order_for_group( fn pool_key_candidate_order_for_group(
group: &EligibleLocalExecutionCandidate, group: &EligibleLocalExecutionCandidate,
routing_policy: Option<&ResolvedRoutingPolicy>,
) -> StoredPoolKeyCandidateOrder { ) -> StoredPoolKeyCandidateOrder {
let Some(pool_config) = pool_config_for_candidate(group) else { let Some(pool_config) = pool_config_for_candidate(group) else {
return StoredPoolKeyCandidateOrder::InternalPriority; return StoredPoolKeyCandidateOrder::InternalPriority;
}; };
let presets = pool_config let override_presets = routing_policy
.scheduling_presets .and_then(|policy| {
.iter() policy
.map(|preset| PoolSchedulingPreset { .pool_policy_overrides
preset: preset.preset.clone(), .get(group.candidate.provider_id.as_str())
enabled: preset.enabled,
mode: preset.mode.clone(),
}) })
.collect::<Vec<_>>(); .filter(|override_policy| !override_policy.scheduling_presets.is_empty());
let presets = match override_presets {
Some(override_policy) => override_policy
.scheduling_presets
.iter()
.map(|preset| PoolSchedulingPreset {
preset: preset.preset.clone(),
enabled: preset.enabled,
mode: preset.mode.clone(),
})
.collect::<Vec<_>>(),
None => pool_config
.scheduling_presets
.iter()
.map(|preset| PoolSchedulingPreset {
preset: preset.preset.clone(),
enabled: preset.enabled,
mode: preset.mode.clone(),
})
.collect::<Vec<_>>(),
};
let active_presets = ProviderPoolService::with_builtin_adapters() let active_presets = ProviderPoolService::with_builtin_adapters()
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets) .normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets)
.into_iter() .into_iter()
@@ -1071,6 +1138,7 @@ mod tests {
apply_local_execution_pool_scheduler_with_runtime_map, build_pool_catalog_key_context, apply_local_execution_pool_scheduler_with_runtime_map, build_pool_catalog_key_context,
pool_config_for_candidate, should_trigger_active_probe_burst_for_request, pool_config_for_candidate, should_trigger_active_probe_burst_for_request,
PoolCatalogKeyContext, PoolKeyCursor, POOL_ACTIVE_PROBE_SEALED_SKIP_REASON, PoolCatalogKeyContext, PoolKeyCursor, POOL_ACTIVE_PROBE_SEALED_SKIP_REASON,
ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON,
}; };
use crate::ai_serving::{ use crate::ai_serving::{
apply_local_runtime_candidate_terminal_reason, EligibleLocalExecutionCandidate, apply_local_runtime_candidate_terminal_reason, EligibleLocalExecutionCandidate,
@@ -1096,6 +1164,9 @@ mod tests {
GatewayProviderTransportEndpoint, GatewayProviderTransportKey, GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
GatewayProviderTransportProvider, GatewayProviderTransportProvider,
}; };
use aether_routing_core::{
RankingOverlay, ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode,
};
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate; use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use serde_json::json; use serde_json::json;
use std::collections::{BTreeMap, BTreeSet, VecDeque}; use std::collections::{BTreeMap, BTreeSet, VecDeque};
@@ -2103,6 +2174,59 @@ mod tests {
); );
} }
#[tokio::test]
async fn pool_key_cursor_filters_expanded_keys_by_routing_profile_allowed_keys() {
let app = AppState::new().expect("state should build");
let provider_config = Some(json!({ "pool_advanced": { "lru_enabled": true } }));
let group = sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"pool-group",
10,
provider_config.clone(),
);
let routing_policy = routing_policy_with_allowed_keys(["key-b"]);
let mut cursor = PoolKeyCursor::new_with_routing_policy(
PlannerAppState::new(&app),
group,
None,
None,
None,
Some(&routing_policy),
);
cursor.queued_candidates = VecDeque::from([
sample_eligible_candidate(
"provider-pool",
"endpoint-1",
"key-a",
10,
provider_config.clone(),
),
sample_eligible_candidate("provider-pool", "endpoint-1", "key-b", 10, provider_config),
]);
let candidate = cursor
.next_key()
.await
.expect("cursor should skip disallowed pool key and return allowed key");
assert_eq!(candidate.candidate.key_id, "key-b");
assert_eq!(candidate.orchestration.pool_key_index, Some(0));
assert_eq!(
cursor
.skip_reason_counts
.get(ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON),
Some(&1)
);
let skipped = cursor.take_skipped_candidates();
assert_eq!(
skipped
.iter()
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
.collect::<Vec<_>>(),
vec![("key-a", ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON)]
);
}
#[tokio::test] #[tokio::test]
async fn pool_key_cursor_allows_parallel_requests_to_use_same_healthy_key() { async fn pool_key_cursor_allows_parallel_requests_to_use_same_healthy_key() {
let app = AppState::new().expect("state should build"); let app = AppState::new().expect("state should build");
@@ -2668,6 +2792,28 @@ mod tests {
(provider, endpoint, keys, rows) (provider, endpoint, keys, rows)
} }
fn routing_policy_with_allowed_keys<const N: usize>(
key_ids: [&str; N],
) -> ResolvedRoutingPolicy {
ResolvedRoutingPolicy {
group_id: Some("routing-group-1".to_string()),
group_version: Some(1),
selection_source: "test".to_string(),
requested_model: "gpt-5".to_string(),
resolved_model: "gpt-5".to_string(),
priority_mode: RoutingSetPriorityMode::Provider,
scheduling_mode: RoutingSchedulingMode::CacheAffinity,
keep_priority_on_conversion: false,
ranking_overlay: RankingOverlay {
allowed_keys: key_ids.into_iter().map(str::to_string).collect(),
..RankingOverlay::default()
},
mutation_plan: Default::default(),
pool_policy_overrides: BTreeMap::new(),
matched_rules: Vec::new(),
}
}
fn sample_eligible_candidate( fn sample_eligible_candidate(
provider_id: &str, provider_id: &str,
endpoint_id: &str, endpoint_id: &str,

View File

@@ -3,9 +3,10 @@ use std::io::Error as IoError;
use std::time::Instant; use std::time::Instant;
use aether_contracts::{ use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResponseBody, StreamFrame, ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResolvedTransportProfile,
StreamFramePayload, StreamFrameType, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, ResponseBody, StreamFrame, StreamFramePayload, StreamFrameType,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
}; };
use axum::body::Bytes; use axum::body::Bytes;
use base64::Engine as _; use base64::Engine as _;
@@ -30,6 +31,7 @@ const CHATGPT_WEB_CLIENT_VERSION: &str = "prod-be885abbfcfe7b1f511e88b3003d9ee44
const CHATGPT_WEB_BUILD_NUMBER: &str = "5955942"; const CHATGPT_WEB_BUILD_NUMBER: &str = "5955942";
const CHATGPT_WEB_SEC_CH_UA: &str = const CHATGPT_WEB_SEC_CH_UA: &str =
r#""Microsoft Edge";v="143", "Chromium";v="143", "Not A(Brand";v="24""#; r#""Microsoft Edge";v="143", "Chromium";v="143", "Not A(Brand";v="24""#;
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
pub(crate) struct ChatGptWebImageStream { pub(crate) struct ChatGptWebImageStream {
pub(crate) frame_stream: BoxStream<'static, Result<Bytes, IoError>>, pub(crate) frame_stream: BoxStream<'static, Result<Bytes, IoError>>,
@@ -921,7 +923,7 @@ async fn execute_subrequest(
provider_api_format: plan.provider_api_format.clone(), provider_api_format: plan.provider_api_format.clone(),
model_name: plan.model_name.clone(), model_name: plan.model_name.clone(),
proxy: plan.proxy.clone(), proxy: plan.proxy.clone(),
transport_profile: plan.transport_profile.clone(), transport_profile: chatgpt_web_image_transport_profile(plan),
timeouts: plan.timeouts.clone(), timeouts: plan.timeouts.clone(),
}; };
DirectSyncExecutionRuntime::new() DirectSyncExecutionRuntime::new()
@@ -929,6 +931,34 @@ async fn execute_subrequest(
.await .await
} }
fn chatgpt_web_image_transport_profile(plan: &ExecutionPlan) -> Option<ResolvedTransportProfile> {
match plan.transport_profile.as_ref() {
Some(profile)
if profile
.backend
.trim()
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
{
Some(profile.clone())
}
_ => Some(default_chatgpt_web_image_transport_profile()),
}
}
fn default_chatgpt_web_image_transport_profile() -> ResolvedTransportProfile {
ResolvedTransportProfile {
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
header_fingerprint: None,
extra: Some(json!({
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
"source": "chatgpt_web_image_default",
})),
}
}
fn web_base_headers(fp: &WebFingerprint, token: &str, path: &str) -> BTreeMap<String, String> { fn web_base_headers(fp: &WebFingerprint, token: &str, path: &str) -> BTreeMap<String, String> {
let mut headers = BTreeMap::from([ let mut headers = BTreeMap::from([
("user-agent".to_string(), fp.user_agent.to_string()), ("user-agent".to_string(), fp.user_agent.to_string()),
@@ -1962,6 +1992,30 @@ mod tests {
} }
} }
#[test]
fn chatgpt_web_image_subrequests_default_to_browser_wreq_transport() {
let plan = sample_plan(
CHATGPT_WEB_DEFAULT_BASE_URL,
json!({"prompt": "draw a small test image"}),
false,
);
let profile = chatgpt_web_image_transport_profile(&plan).expect("transport profile");
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
assert_eq!(
profile
.extra
.as_ref()
.and_then(|value| value.get("source"))
.and_then(Value::as_str),
Some("chatgpt_web_image_default")
);
}
async fn start_mock_chatgpt_web() -> (String, tokio::task::JoinHandle<()>) { async fn start_mock_chatgpt_web() -> (String, tokio::task::JoinHandle<()>) {
let app = Router::new().fallback(any(|request: Request| async move { let app = Router::new().fallback(any(|request: Request| async move {
let path = request.uri().path().to_string(); let path = request.uri().path().to_string();

File diff suppressed because it is too large Load Diff

View File

@@ -6,6 +6,7 @@ use serde_json::{Map, Value};
mod chatgpt_web_image; mod chatgpt_web_image;
mod constants; mod constants;
mod fallback; mod fallback;
mod grok;
mod kiro_web_search; mod kiro_web_search;
pub(crate) mod ndjson; pub(crate) mod ndjson;
mod oauth_retry; mod oauth_retry;
@@ -54,6 +55,7 @@ pub(crate) use sync::{
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild, resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
LocalVideoSyncSuccessOutcome, LocalVideoSyncSuccessOutcome,
}; };
pub(crate) use transport::execute_sync_plan_with_report_context as execute_execution_runtime_sync_plan_with_report_context;
pub(crate) use transport::{ pub(crate) use transport::{
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime, execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,
DirectUpstreamStreamExecution, ExecutionRuntimeTransportError, DirectUpstreamStreamExecution, ExecutionRuntimeTransportError,

View File

@@ -370,6 +370,8 @@ impl IntoResponse for ExecutionRuntimeAppError {
) => StatusCode::BAD_REQUEST, ) => StatusCode::BAD_REQUEST,
ExecutionRuntimeServerError::Transport( ExecutionRuntimeServerError::Transport(
ExecutionRuntimeTransportError::ClientBuild(_) ExecutionRuntimeTransportError::ClientBuild(_)
| ExecutionRuntimeTransportError::BrowserClientBuild(_)
| ExecutionRuntimeTransportError::BrowserBody(_)
| ExecutionRuntimeTransportError::UpstreamRequest(_) | ExecutionRuntimeTransportError::UpstreamRequest(_)
| ExecutionRuntimeTransportError::RelayError(_) | ExecutionRuntimeTransportError::RelayError(_)
| ExecutionRuntimeTransportError::InvalidJson(_), | ExecutionRuntimeTransportError::InvalidJson(_),

View File

@@ -60,6 +60,7 @@ use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER};
use crate::control::GatewayControlDecision; use crate::control::GatewayControlDecision;
use crate::execution_runtime::build_direct_execution_frame_stream; use crate::execution_runtime::build_direct_execution_frame_stream;
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_stream; use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_stream;
use crate::execution_runtime::grok::maybe_execute_grok_stream;
use crate::execution_runtime::kiro_web_search::maybe_execute_kiro_web_search_stream; use crate::execution_runtime::kiro_web_search::maybe_execute_kiro_web_search_stream;
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry; use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
#[cfg(test)] #[cfg(test)]
@@ -525,6 +526,58 @@ pub(crate) async fn execute_execution_runtime_stream(
key_id.as_str(), key_id.as_str(),
) )
.await; .await;
match maybe_execute_grok_stream(&plan, report_context.as_ref()).await {
Ok(Some(grok_stream)) => {
return execute_stream_from_frame_stream(
state,
plan,
trace_id,
decision,
plan_kind,
report_kind,
grok_stream.report_context.or(report_context),
candidate_started_unix_secs,
stream_started_at,
grok_stream.frame_stream,
provider_pool_in_flight_guard.take(),
)
.await;
}
Ok(None) => {}
Err(err) => {
info!(
event_name = "grok_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan.candidate_id,
provider_name = provider_name.as_str(),
endpoint_id = %endpoint_id,
key_id = %key_id,
model_name = model_name.as_str(),
candidate_index = candidate_index.as_str(),
error = %err,
"gateway Grok stream execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(format!("{err:?}")),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
}
match maybe_execute_kiro_web_search_stream(state, &plan, report_context.as_ref()).await { match maybe_execute_kiro_web_search_stream(state, &plan, report_context.as_ref()).await {
Ok(Some(kiro_web_search)) => { Ok(Some(kiro_web_search)) => {
return execute_stream_from_frame_stream( return execute_stream_from_frame_stream(

View File

@@ -18,7 +18,9 @@ use crate::ai_serving::api::{
normalize_provider_private_report_context, StreamingStandardTerminalObserver, normalize_provider_private_report_context, StreamingStandardTerminalObserver,
}; };
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson; use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
use crate::execution_runtime::transport::DirectUpstreamResponse; use crate::execution_runtime::transport::{
format_wreq_upstream_request_error, DirectUpstreamResponse,
};
use crate::execution_runtime::DirectUpstreamStreamExecution; use crate::execution_runtime::DirectUpstreamStreamExecution;
use crate::GatewayError; use crate::GatewayError;
@@ -235,6 +237,62 @@ pub(crate) fn build_direct_execution_frame_stream(
} }
} }
} }
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await {
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
if !first_chunk_telemetry_emitted {
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
first_chunk_telemetry_emitted = true;
}
upstream_bytes += chunk.len() as u64;
observe_stream_chunk(
&mut stream_terminal_observer,
&normalized_observer_context,
private_stream_normalizer.as_mut(),
&mut observer_buffered,
chunk.as_ref(),
);
match encode_data_frame(&chunk) {
Ok(frame) => yield Ok(frame),
Err(err) => {
yield Err(err);
return;
}
}
}
Err(err) => {
let message = format_wreq_upstream_request_error(&err);
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
status_code,
upstream_bytes,
error = %message,
"upstream body stream read error"
);
match encode_error_frame(status_code, message) {
Ok(frame) => yield Ok(frame),
Err(encode_err) => {
yield Err(encode_err);
return;
}
}
break;
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop { DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await { match response.next_chunk().await {
Ok(Some(chunk)) => { Ok(Some(chunk)) => {
@@ -454,6 +512,35 @@ async fn buffer_non_sse_upstream_body(
} }
} }
} }
DirectUpstreamResponse::BrowserWreq(response) => {
let mut bytes_stream = response.bytes_stream();
while let Some(item) = bytes_stream.next().await {
match item {
Ok(chunk) => {
if ttfb_ms.is_none() {
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
}
upstream_bytes += chunk.len() as u64;
body_bytes.extend_from_slice(&chunk);
}
Err(err) => {
let message = format_wreq_upstream_request_error(&err);
warn!(
event_name = "stream_pump_body_read_error",
log_type = "ops",
upstream_bytes,
error = %message,
"upstream body stream read error"
);
return Err(BufferedUpstreamBodyError {
message,
ttfb_ms,
upstream_bytes,
});
}
}
}
}
DirectUpstreamResponse::LocalTunnel(mut response) => loop { DirectUpstreamResponse::LocalTunnel(mut response) => loop {
match response.next_chunk().await { match response.next_chunk().await {
Ok(Some(chunk)) => { Ok(Some(chunk)) => {

View File

@@ -39,14 +39,15 @@ use crate::api::response::{
use crate::clock::current_unix_ms as current_request_candidate_unix_ms; use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
use crate::control::GatewayControlDecision; use crate::control::GatewayControlDecision;
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync; use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
use crate::execution_runtime::grok::maybe_execute_grok_sync;
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry; use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
#[cfg(test)] #[cfg(test)]
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime; use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
use crate::execution_runtime::submission::submit_local_core_error_or_sync_finalize; use crate::execution_runtime::submission::submit_local_core_error_or_sync_finalize;
use crate::execution_runtime::transport::{ use crate::execution_runtime::transport::{
build_request_body, collect_response_headers, decode_response_body_bytes, build_request_body, collect_response_headers, decode_response_body_bytes,
response_body_is_json, send_request, DirectSyncExecutionRuntime, format_upstream_request_error, format_wreq_upstream_request_error, response_body_is_json,
ExecutionRuntimeTransportError, send_request, DirectHttpResponse, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError,
}; };
use crate::execution_runtime::{ use crate::execution_runtime::{
analyze_local_candidate_failover_sync, apply_endpoint_response_header_rules, analyze_local_candidate_failover_sync, apply_endpoint_response_header_rules,
@@ -734,23 +735,46 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
.await .await
.map_err(SyncExecutionFailure::from_transport)?; .map_err(SyncExecutionFailure::from_transport)?;
let ttfb_ms = started_at.elapsed().as_millis() as u64; let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16(); let status_code = response.status_code();
let headers = collect_response_headers(response.headers()); let headers = response.headers();
progress.record_response_started(status_code, ttfb_ms).await; progress.record_response_started(status_code, ttfb_ms).await;
let mut upstream_stream = response.bytes_stream();
let mut body_bytes = Vec::new(); let mut body_bytes = Vec::new();
while let Some(chunk) = upstream_stream.next().await { match response {
let chunk = chunk.map_err(|err| { DirectHttpResponse::Reqwest(response) => {
SyncExecutionFailure::from_transport(ExecutionRuntimeTransportError::UpstreamRequest( let mut upstream_stream = response.bytes_stream();
crate::execution_runtime::transport::format_upstream_request_error(&err), while let Some(chunk) = upstream_stream.next().await {
)) let chunk = chunk.map_err(|err| {
})?; SyncExecutionFailure::from_transport(
let elapsed_ms = started_at.elapsed().as_millis() as u64; ExecutionRuntimeTransportError::UpstreamRequest(
progress format_upstream_request_error(&err),
.observe_chunk(&chunk, status_code, elapsed_ms) ),
.await; )
body_bytes.extend_from_slice(&chunk); })?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
.await;
body_bytes.extend_from_slice(&chunk);
}
}
DirectHttpResponse::BrowserWreq(response) => {
let mut upstream_stream = response.bytes_stream();
while let Some(chunk) = upstream_stream.next().await {
let chunk = chunk.map_err(|err| {
SyncExecutionFailure::from_transport(
ExecutionRuntimeTransportError::UpstreamRequest(
format_wreq_upstream_request_error(&err),
),
)
})?;
let elapsed_ms = started_at.elapsed().as_millis() as u64;
progress
.observe_chunk(&chunk, status_code, elapsed_ms)
.await;
body_bytes.extend_from_slice(&chunk);
}
}
} }
let decoded_body_bytes = let decoded_body_bytes =
@@ -1129,64 +1153,106 @@ async fn execute_execution_runtime_sync_impl(
.await; .await;
#[cfg(not(test))] #[cfg(not(test))]
let mut result = { let mut result = {
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await { match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
Ok(Some(result)) => result, Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate( Ok(None) => {
state, match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref())
&plan, .await
report_context.as_ref(), {
trace_id, Ok(Some(result)) => result,
plan_kind, Ok(None) => match execute_direct_sync_runtime_candidate(
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await
{
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state, state,
&plan, &plan,
report_context.as_ref(), report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate { trace_id,
status: RequestCandidateStatus::Failed, plan_kind,
status_code: err.status_code, plan_request_id_for_log.as_str(),
error_type: Some(err.error_type.to_string()), plan_candidate_id.as_deref(),
error_message: Some(err.message), provider_name.as_str(),
latency_ms: err.latency_ms, endpoint_id.as_str(),
started_at_unix_ms: Some(candidate_started_unix_secs), key_id.as_str(),
finished_at_unix_ms: Some(terminal_unix_secs), model_name.as_str(),
}, candidate_index.as_str(),
progress_snapshot.clone(),
) )
.await; .await
return Ok(None); {
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
},
Err(err) => {
warn!(
event_name = "chatgpt_web_image_execution_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error = %err,
"gateway ChatGPT-Web image execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: None,
error_type: Some(
"chatgpt_web_image_execution_unavailable".to_string(),
),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
} }
}, }
Err(err) => { Err(err) => {
warn!( warn!(
event_name = "chatgpt_web_image_execution_unavailable", event_name = "grok_execution_unavailable",
log_type = "ops", log_type = "ops",
trace_id = %trace_id, trace_id = %trace_id,
request_id = %plan_request_id_for_log, request_id = %plan_request_id_for_log,
@@ -1197,7 +1263,7 @@ async fn execute_execution_runtime_sync_impl(
model_name, model_name,
candidate_index = candidate_index.as_str(), candidate_index = candidate_index.as_str(),
error = %err, error = %err,
"gateway ChatGPT-Web image execution unavailable" "gateway Grok execution unavailable"
); );
let terminal_unix_secs = current_request_candidate_unix_ms(); let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status( record_local_request_candidate_status(
@@ -1207,7 +1273,7 @@ async fn execute_execution_runtime_sync_impl(
SchedulerRequestCandidateStatusUpdate { SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed, status: RequestCandidateStatus::Failed,
status_code: None, status_code: None,
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()), error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(err.to_string()), error_message: Some(err.to_string()),
latency_ms: None, latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs), started_at_unix_ms: Some(candidate_started_unix_secs),
@@ -1264,30 +1330,72 @@ async fn execute_execution_runtime_sync_impl(
.trim() .trim()
.is_empty() .is_empty()
{ {
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
{
Ok(Some(result)) => result, Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate( Ok(None) => match maybe_execute_chatgpt_web_image_sync(
state, state,
&plan, &plan,
report_context.as_ref(), report_context.as_ref(),
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
) )
.await .await
{ {
Ok(result) => result, Ok(Some(result)) => result,
Ok(None) => match execute_direct_sync_runtime_candidate(
state,
&plan,
report_context.as_ref(),
trace_id,
plan_kind,
plan_request_id_for_log.as_str(),
plan_candidate_id.as_deref(),
provider_name.as_str(),
endpoint_id.as_str(),
key_id.as_str(),
model_name.as_str(),
candidate_index.as_str(),
progress_snapshot.clone(),
)
.await
{
Ok(result) => result,
Err(err) => {
warn!(
event_name = "sync_execution_runtime_unavailable",
log_type = "ops",
trace_id = %trace_id,
request_id = %plan_request_id_for_log,
candidate_id = ?plan_candidate_id,
provider_name,
endpoint_id,
key_id,
model_name,
candidate_index = candidate_index.as_str(),
error_type = err.error_type,
error = %err.message,
"gateway in-process sync execution unavailable"
);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: err.status_code,
error_type: Some(err.error_type.to_string()),
error_message: Some(err.message),
latency_ms: err.latency_ms,
started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs),
},
)
.await;
return Ok(None);
}
},
Err(err) => { Err(err) => {
warn!( warn!(
event_name = "sync_execution_runtime_unavailable", event_name = "chatgpt_web_image_execution_unavailable",
log_type = "ops", log_type = "ops",
trace_id = %trace_id, trace_id = %trace_id,
request_id = %plan_request_id_for_log, request_id = %plan_request_id_for_log,
@@ -1297,9 +1405,8 @@ async fn execute_execution_runtime_sync_impl(
key_id, key_id,
model_name, model_name,
candidate_index = candidate_index.as_str(), candidate_index = candidate_index.as_str(),
error_type = err.error_type, error = %err,
error = %err.message, "gateway ChatGPT-Web image execution unavailable"
"gateway in-process sync execution unavailable"
); );
let terminal_unix_secs = current_request_candidate_unix_ms(); let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status( record_local_request_candidate_status(
@@ -1308,10 +1415,12 @@ async fn execute_execution_runtime_sync_impl(
report_context.as_ref(), report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate { SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed, status: RequestCandidateStatus::Failed,
status_code: err.status_code, status_code: None,
error_type: Some(err.error_type.to_string()), error_type: Some(
error_message: Some(err.message), "chatgpt_web_image_execution_unavailable".to_string(),
latency_ms: err.latency_ms, ),
error_message: Some(err.to_string()),
latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs), started_at_unix_ms: Some(candidate_started_unix_secs),
finished_at_unix_ms: Some(terminal_unix_secs), finished_at_unix_ms: Some(terminal_unix_secs),
}, },
@@ -1322,7 +1431,7 @@ async fn execute_execution_runtime_sync_impl(
}, },
Err(err) => { Err(err) => {
warn!( warn!(
event_name = "chatgpt_web_image_execution_unavailable", event_name = "grok_execution_unavailable",
log_type = "ops", log_type = "ops",
trace_id = %trace_id, trace_id = %trace_id,
request_id = %plan_request_id_for_log, request_id = %plan_request_id_for_log,
@@ -1333,7 +1442,7 @@ async fn execute_execution_runtime_sync_impl(
model_name, model_name,
candidate_index = candidate_index.as_str(), candidate_index = candidate_index.as_str(),
error = %err, error = %err,
"gateway ChatGPT-Web image execution unavailable" "gateway Grok execution unavailable"
); );
let terminal_unix_secs = current_request_candidate_unix_ms(); let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status( record_local_request_candidate_status(
@@ -1343,7 +1452,7 @@ async fn execute_execution_runtime_sync_impl(
SchedulerRequestCandidateStatusUpdate { SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed, status: RequestCandidateStatus::Failed,
status_code: None, status_code: None,
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()), error_type: Some("grok_execution_unavailable".to_string()),
error_message: Some(err.to_string()), error_message: Some(err.to_string()),
latency_ms: None, latency_ms: None,
started_at_unix_ms: Some(candidate_started_unix_secs), started_at_unix_ms: Some(candidate_started_unix_secs),

View File

@@ -8,7 +8,8 @@ use aether_contracts::{
ExecutionPlan, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile, ExecutionPlan, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile,
ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_HTTP1_ONLY, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
}; };
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation; use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
use aether_http::{apply_http_client_config, HttpClientConfig}; use aether_http::{apply_http_client_config, HttpClientConfig};
@@ -81,6 +82,52 @@ pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
detail detail
} }
pub(crate) fn format_wreq_upstream_request_error(err: &wreq::Error) -> String {
let mut kinds = Vec::new();
if err.is_connect() {
kinds.push("connect");
}
if err.is_timeout() {
kinds.push("timeout");
}
if err.is_redirect() {
kinds.push("redirect");
}
if err.is_body() {
kinds.push("body");
}
if err.is_decode() {
kinds.push("decode");
}
if err.is_request() {
kinds.push("request");
}
let mut detail = err.to_string();
let mut source = err.source();
while let Some(cause) = source {
let cause_text = cause.to_string();
if !cause_text.is_empty() && !detail.contains(&cause_text) {
detail.push_str(": ");
detail.push_str(&cause_text);
}
source = cause.source();
}
if let Some(uri) = err.uri() {
detail.push_str(" [uri=");
detail.push_str(&uri.to_string());
detail.push(']');
}
if !kinds.is_empty() {
detail.push_str(" [kind=");
detail.push_str(&kinds.join(","));
detail.push(']');
}
detail
}
#[derive(Debug, Error)] #[derive(Debug, Error)]
pub(crate) enum ExecutionRuntimeTransportError { pub(crate) enum ExecutionRuntimeTransportError {
#[error("stream execution is not supported for this plan")] #[error("stream execution is not supported for this plan")]
@@ -107,6 +154,10 @@ pub(crate) enum ExecutionRuntimeTransportError {
BodyEncode(serde_json::Error), BodyEncode(serde_json::Error),
#[error("failed to build HTTP client: {0}")] #[error("failed to build HTTP client: {0}")]
ClientBuild(reqwest::Error), ClientBuild(reqwest::Error),
#[error("failed to build browser impersonation HTTP client: {0}")]
BrowserClientBuild(wreq::Error),
#[error("browser impersonation response body failed: {0}")]
BrowserBody(String),
#[error("failed to execute upstream request: {0}")] #[error("failed to execute upstream request: {0}")]
UpstreamRequest(String), UpstreamRequest(String),
#[error("hub relay request failed: {0}")] #[error("hub relay request failed: {0}")]
@@ -136,7 +187,7 @@ struct RelayRequestMeta {
pub(crate) struct DirectSyncExecutionRuntime; pub(crate) struct DirectSyncExecutionRuntime;
#[derive(Debug, Clone, Copy, Default)] #[derive(Debug, Clone, Copy, Default)]
struct ExecutionTransportControls { pub(crate) struct ExecutionTransportControls {
follow_redirects: Option<bool>, follow_redirects: Option<bool>,
http1_only: bool, http1_only: bool,
accept_invalid_certs: bool, accept_invalid_certs: bool,
@@ -144,6 +195,7 @@ struct ExecutionTransportControls {
pub(crate) enum DirectUpstreamResponse { pub(crate) enum DirectUpstreamResponse {
Reqwest(reqwest::Response), Reqwest(reqwest::Response),
BrowserWreq(wreq::Response),
LocalTunnel(tunnel::DirectRelayResponse), LocalTunnel(tunnel::DirectRelayResponse),
} }
@@ -172,11 +224,9 @@ impl DirectSyncExecutionRuntime {
let started_at = Instant::now(); let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?; let response = send_request(plan, body_bytes).await?;
let ttfb_ms = started_at.elapsed().as_millis() as u64; let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16(); let status_code = response.status_code();
let headers = collect_response_headers(response.headers()); let headers = response.headers();
let body_bytes = response.bytes().await.map_err(|err| { let body_bytes = response.bytes().await?;
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})?;
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes) let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
.unwrap_or_else(|| body_bytes.to_vec()); .unwrap_or_else(|| body_bytes.to_vec());
let elapsed_ms = started_at.elapsed().as_millis() as u64; let elapsed_ms = started_at.elapsed().as_millis() as u64;
@@ -230,8 +280,8 @@ impl DirectSyncExecutionRuntime {
let started_at = Instant::now(); let started_at = Instant::now();
let response = send_request(plan, body_bytes).await?; let response = send_request(plan, body_bytes).await?;
let status_code = response.status().as_u16(); let status_code = response.status_code();
let headers = collect_response_headers(response.headers()); let headers = response.headers();
let stream_summary_report_context = build_stream_summary_report_context(plan); let stream_summary_report_context = build_stream_summary_report_context(plan);
@@ -242,7 +292,7 @@ impl DirectSyncExecutionRuntime {
headers, headers,
provider_api_format: plan.provider_api_format.clone(), provider_api_format: plan.provider_api_format.clone(),
stream_summary_report_context, stream_summary_report_context,
response: DirectUpstreamResponse::Reqwest(response), response: response.into_direct_upstream_response(),
started_at, started_at,
}) })
} }
@@ -252,6 +302,15 @@ pub(crate) async fn execute_sync_plan(
state: &AppState, state: &AppState,
trace_id: Option<&str>, trace_id: Option<&str>,
plan: &ExecutionPlan, plan: &ExecutionPlan,
) -> Result<ExecutionResult, GatewayError> {
execute_sync_plan_with_report_context(state, trace_id, plan, None).await
}
pub(crate) async fn execute_sync_plan_with_report_context(
state: &AppState,
trace_id: Option<&str>,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
) -> Result<ExecutionResult, GatewayError> { ) -> Result<ExecutionResult, GatewayError> {
#[cfg(test)] #[cfg(test)]
{ {
@@ -275,6 +334,18 @@ pub(crate) async fn execute_sync_plan(
.map_err(|err| GatewayError::Internal(err.to_string())); .map_err(|err| GatewayError::Internal(err.to_string()));
} }
match super::grok::maybe_execute_grok_sync(plan, report_context).await {
Ok(Some(result)) => {
record_manual_proxy_request_outcome(state, plan, result.status_code).await;
return Ok(result);
}
Ok(None) => {}
Err(err) => {
record_manual_proxy_request_failure(state, plan).await;
return Err(GatewayError::Internal(err.to_string()));
}
}
let _ = trace_id; let _ = trace_id;
match DirectSyncExecutionRuntime::new().execute_sync(plan).await { match DirectSyncExecutionRuntime::new().execute_sync(plan).await {
Ok(result) => { Ok(result) => {
@@ -554,7 +625,7 @@ fn build_direct_tunnel_request_meta(
pub(crate) async fn send_request( pub(crate) async fn send_request(
plan: &ExecutionPlan, plan: &ExecutionPlan,
body_bytes: Vec<u8>, body_bytes: Vec<u8>,
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> { ) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) { if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) {
return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail)); return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail));
} }
@@ -572,6 +643,18 @@ pub(crate) async fn send_request(
.and_then(|timeouts| timeouts.total_ms) .and_then(|timeouts| timeouts.total_ms)
.map(Duration::from_millis); .map(Duration::from_millis);
if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) {
return send_via_browser_wreq_transport(
plan,
method,
headers,
body_bytes,
total_timeout,
transport_controls,
)
.await;
}
if let Some(node_id) = resolve_tunnel_node_id(plan.proxy.as_ref()) { if let Some(node_id) = resolve_tunnel_node_id(plan.proxy.as_ref()) {
return send_via_tunnel_relay( return send_via_tunnel_relay(
plan, plan,
@@ -582,7 +665,8 @@ pub(crate) async fn send_request(
total_timeout, total_timeout,
transport_controls, transport_controls,
) )
.await; .await
.map(DirectHttpResponse::Reqwest);
} }
let client = build_client( let client = build_client(
@@ -596,9 +680,95 @@ pub(crate) async fn send_request(
if let Some(timeout) = total_timeout { if let Some(timeout) = total_timeout {
request = request.timeout(timeout); request = request.timeout(timeout);
} }
request.send().await.map_err(|err| { request
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err)) .send()
}) .await
.map(DirectHttpResponse::Reqwest)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
})
}
pub(crate) enum DirectHttpResponse {
Reqwest(reqwest::Response),
BrowserWreq(wreq::Response),
}
impl DirectHttpResponse {
pub(crate) fn status_code(&self) -> u16 {
match self {
DirectHttpResponse::Reqwest(response) => response.status().as_u16(),
DirectHttpResponse::BrowserWreq(response) => response.status().as_u16(),
}
}
pub(crate) fn headers(&self) -> BTreeMap<String, String> {
match self {
DirectHttpResponse::Reqwest(response) => collect_response_headers(response.headers()),
DirectHttpResponse::BrowserWreq(response) => {
collect_response_headers(response.headers())
}
}
}
pub(crate) async fn bytes(self) -> Result<Bytes, ExecutionRuntimeTransportError> {
match self {
DirectHttpResponse::Reqwest(response) => response.bytes().await.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
}),
DirectHttpResponse::BrowserWreq(response) => response.bytes().await.map_err(|err| {
ExecutionRuntimeTransportError::BrowserBody(format_wreq_upstream_request_error(
&err,
))
}),
}
}
fn into_direct_upstream_response(self) -> DirectUpstreamResponse {
match self {
DirectHttpResponse::Reqwest(response) => DirectUpstreamResponse::Reqwest(response),
DirectHttpResponse::BrowserWreq(response) => {
DirectUpstreamResponse::BrowserWreq(response)
}
}
}
}
async fn send_via_browser_wreq_transport(
plan: &ExecutionPlan,
method: reqwest::Method,
headers: HeaderMap,
body_bytes: Vec<u8>,
total_timeout: Option<Duration>,
transport_controls: ExecutionTransportControls,
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
let profile = plan.transport_profile.as_ref().ok_or_else(|| {
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
})?;
let client = build_browser_wreq_client(
plan.timeouts.as_ref(),
plan.proxy.as_ref(),
profile,
transport_controls,
)?;
let method = wreq::Method::from_bytes(method.as_str().as_bytes())
.map_err(ExecutionRuntimeTransportError::InvalidMethod)?;
let mut request = client
.request(method, plan.url.as_str())
.headers(headers)
.body(body_bytes);
if let Some(timeout) = total_timeout {
request = request.timeout(timeout);
}
request
.send()
.await
.map(DirectHttpResponse::BrowserWreq)
.map_err(|err| {
ExecutionRuntimeTransportError::UpstreamRequest(format_wreq_upstream_request_error(
&err,
))
})
} }
async fn send_via_tunnel_relay( async fn send_via_tunnel_relay(
@@ -905,6 +1075,96 @@ fn build_client(
.map_err(ExecutionRuntimeTransportError::ClientBuild) .map_err(ExecutionRuntimeTransportError::ClientBuild)
} }
pub(crate) fn build_browser_wreq_client(
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
proxy: Option<&ProxySnapshot>,
transport_profile: &ResolvedTransportProfile,
transport_controls: ExecutionTransportControls,
) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
let mut builder = wreq::Client::builder().emulation(emulation);
if transport_controls.follow_redirects == Some(true) {
builder = builder.redirect(wreq::redirect::Policy::limited(10));
}
if transport_controls.http1_only || transport_profile_http1_only(Some(transport_profile)) {
builder = builder.http1_only();
}
if transport_controls.accept_invalid_certs {
builder = builder.cert_verification(false).verify_hostname(false);
}
if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) {
builder = builder.connect_timeout(Duration::from_millis(connect_ms));
}
if let Some(total_ms) = timeouts.and_then(|timeouts| timeouts.total_ms) {
builder = builder.timeout(Duration::from_millis(total_ms));
}
if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) {
builder = builder.read_timeout(Duration::from_millis(read_ms));
}
if let Some(proxy_url) = resolve_proxy_url(proxy)? {
let proxy = wreq::Proxy::all(proxy_url.as_str())
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)?;
builder = builder.proxy(proxy);
}
builder
.build()
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)
}
fn browser_wreq_emulation_from_profile(
profile: &ResolvedTransportProfile,
) -> Result<wreq_util::Emulation, ExecutionRuntimeTransportError> {
match normalize_browser_profile_name(browser_transport_profile_name(profile)).as_str() {
"chrome100" => Ok(wreq_util::Emulation::Chrome100),
"chrome101" => Ok(wreq_util::Emulation::Chrome101),
"chrome104" => Ok(wreq_util::Emulation::Chrome104),
"chrome105" => Ok(wreq_util::Emulation::Chrome105),
"chrome106" => Ok(wreq_util::Emulation::Chrome106),
"chrome107" => Ok(wreq_util::Emulation::Chrome107),
"chrome108" => Ok(wreq_util::Emulation::Chrome108),
"chrome109" => Ok(wreq_util::Emulation::Chrome109),
"chrome110" => Ok(wreq_util::Emulation::Chrome110),
"chrome114" => Ok(wreq_util::Emulation::Chrome114),
"chrome116" => Ok(wreq_util::Emulation::Chrome116),
"chrome117" => Ok(wreq_util::Emulation::Chrome117),
"chrome118" => Ok(wreq_util::Emulation::Chrome118),
"chrome119" => Ok(wreq_util::Emulation::Chrome119),
"chrome120" => Ok(wreq_util::Emulation::Chrome120),
"chrome123" => Ok(wreq_util::Emulation::Chrome123),
"chrome124" => Ok(wreq_util::Emulation::Chrome124),
"chrome126" => Ok(wreq_util::Emulation::Chrome126),
"chrome127" => Ok(wreq_util::Emulation::Chrome127),
"chrome128" => Ok(wreq_util::Emulation::Chrome128),
"chrome129" => Ok(wreq_util::Emulation::Chrome129),
"chrome130" => Ok(wreq_util::Emulation::Chrome130),
"chrome131" => Ok(wreq_util::Emulation::Chrome131),
"chrome132" => Ok(wreq_util::Emulation::Chrome132),
"chrome133" => Ok(wreq_util::Emulation::Chrome133),
"chrome134" => Ok(wreq_util::Emulation::Chrome134),
"chrome135" => Ok(wreq_util::Emulation::Chrome135),
"chrome136" => Ok(wreq_util::Emulation::Chrome136),
"chrome137" => Ok(wreq_util::Emulation::Chrome137),
"chrome138" => Ok(wreq_util::Emulation::Chrome138),
"chrome139" => Ok(wreq_util::Emulation::Chrome139),
"chrome140" => Ok(wreq_util::Emulation::Chrome140),
"chrome141" => Ok(wreq_util::Emulation::Chrome141),
"chrome142" => Ok(wreq_util::Emulation::Chrome142),
"chrome143" => Ok(wreq_util::Emulation::Chrome143),
"chrome144" => Ok(wreq_util::Emulation::Chrome144),
"chrome145" => Ok(wreq_util::Emulation::Chrome145),
other => Err(ExecutionRuntimeTransportError::UnsupportedTransportProfile(
format!("browser_wreq:{other}"),
)),
}
}
fn normalize_browser_profile_name(value: String) -> String {
value
.trim()
.to_ascii_lowercase()
.replace(['_', '-', ' '], "")
}
fn validate_reqwest_transport_profile( fn validate_reqwest_transport_profile(
transport_profile: Option<&ResolvedTransportProfile>, transport_profile: Option<&ResolvedTransportProfile>,
) -> Result<(), ExecutionRuntimeTransportError> { ) -> Result<(), ExecutionRuntimeTransportError> {
@@ -923,6 +1183,56 @@ fn validate_reqwest_transport_profile(
)) ))
} }
fn transport_profile_uses_browser_wreq(
transport_profile: Option<&ResolvedTransportProfile>,
) -> bool {
transport_profile
.map(|profile| {
profile
.backend
.trim()
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ)
})
.unwrap_or(false)
}
fn browser_transport_profile_name(profile: &ResolvedTransportProfile) -> String {
profile
.extra
.as_ref()
.and_then(|value| {
value
.get("browser_profile")
.or_else(|| value.get("impersonate"))
.and_then(Value::as_str)
})
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
profile
.profile_id
.trim()
.is_empty()
.then_some("chrome136".to_string())
.or_else(|| Some(profile.profile_id.trim().to_string()))
})
.unwrap_or_else(|| "chrome136".to_string())
}
fn insert_browser_control_header(
headers: &mut HeaderMap,
name: &'static str,
value: &str,
) -> Result<(), ExecutionRuntimeTransportError> {
headers.insert(
HeaderName::from_static(name),
HeaderValue::from_str(value)
.map_err(|_| ExecutionRuntimeTransportError::InvalidHeaderValue(name.to_string()))?,
);
Ok(())
}
fn transport_profile_http1_only(transport_profile: Option<&ResolvedTransportProfile>) -> bool { fn transport_profile_http1_only(transport_profile: Option<&ResolvedTransportProfile>) -> bool {
transport_profile transport_profile
.map(|profile| { .map(|profile| {
@@ -991,7 +1301,7 @@ fn resolve_proxy_url(
Ok(None) Ok(None)
} }
fn build_request_headers( pub(crate) fn build_request_headers(
headers: &BTreeMap<String, String>, headers: &BTreeMap<String, String>,
content_encoding: Option<&str>, content_encoding: Option<&str>,
allow_passthrough_content_encoding: bool, allow_passthrough_content_encoding: bool,
@@ -1180,22 +1490,22 @@ mod tests {
use aether_contracts::{ use aether_contracts::{
ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile, ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile,
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
}; };
use aether_data::repository::proxy_nodes::{ use aether_data::repository::proxy_nodes::{
InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode, InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode,
}; };
use axum::body::Bytes; use axum::body::{Body, Bytes};
use axum::extract::ws::Message; use axum::extract::ws::Message;
use axum::extract::Path; use axum::extract::Path;
use axum::http::HeaderMap as AxumHeaderMap; use axum::http::HeaderMap as AxumHeaderMap;
use axum::routing::post; use axum::routing::{any, post};
use axum::{Json, Router}; use axum::{Json, Router};
use serde_json::json; use serde_json::json;
use tokio::sync::watch; use tokio::sync::watch;
use super::{ use super::{
build_client, build_request_headers, execute_sync_plan, build_browser_wreq_client, build_client, build_request_headers, execute_sync_plan,
record_manual_proxy_request_failure, record_manual_proxy_request_outcome, record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
record_manual_proxy_request_success, record_manual_proxy_stream_error, record_manual_proxy_request_success, record_manual_proxy_stream_error,
resolve_execution_transport_controls, DirectSyncExecutionRuntime, resolve_execution_transport_controls, DirectSyncExecutionRuntime,
@@ -1431,6 +1741,228 @@ mod tests {
); );
} }
#[tokio::test]
async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() {
async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response {
assert_eq!(
headers
.get("content-type")
.and_then(|value| value.to_str().ok()),
Some("application/json")
);
assert!(
headers
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
.is_none(),
"internal execution control headers must not leak upstream"
);
assert_eq!(body.as_ref(), br#"{"modelName":"auto"}"#);
axum::response::Response::builder()
.status(http::StatusCode::ACCEPTED)
.header("content-type", "application/json")
.body(Body::from(
json!({
"ok": true,
"via": "browser_wreq"
})
.to_string(),
))
.expect("response should build")
}
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route("/request", any(browser_upstream));
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let plan = ExecutionPlan {
request_id: "req-browser-wreq".into(),
candidate_id: None,
provider_name: Some("grok".into()),
provider_id: "provider-1".into(),
endpoint_id: "endpoint-1".into(),
key_id: "key-1".into(),
method: "POST".into(),
url: format!("http://{addr}/request"),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
(
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.into(),
"true".into(),
),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({"modelName":"auto"})),
stream: false,
client_api_format: "openai:responses".into(),
provider_api_format: "grok:rate_limits".into(),
model_name: Some("grok-quota".into()),
proxy: None,
transport_profile: Some(ResolvedTransportProfile {
profile_id: "chrome136".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: "auto".into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: Some(json!({
"browser_profile": "chrome136"
})),
}),
timeouts: Some(ExecutionTimeouts {
total_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
};
let result = DirectSyncExecutionRuntime::new()
.execute_sync(&plan)
.await
.expect("browser wreq transport plan should execute in-process");
server.abort();
assert_eq!(result.status_code, http::StatusCode::ACCEPTED.as_u16());
assert_eq!(
result
.body
.and_then(|body| body.json_body)
.and_then(|body| body.get("via").cloned()),
Some(json!("browser_wreq"))
);
}
#[test]
fn browser_wreq_transport_rejects_unknown_profile() {
let profile = ResolvedTransportProfile {
profile_id: "firefox999".into(),
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
http_mode: "auto".into(),
pool_scope: "key".into(),
header_fingerprint: None,
extra: None,
};
let error = match build_browser_wreq_client(
None,
None,
&profile,
ExecutionTransportControls::default(),
) {
Ok(_) => panic!("unknown browser profile should fail loudly"),
Err(error) => error,
};
assert!(matches!(
error,
ExecutionRuntimeTransportError::UnsupportedTransportProfile(backend)
if backend == "browser_wreq:firefox999"
));
}
#[tokio::test]
async fn execute_sync_plan_routes_grok_marker_through_grok_runtime() {
let listener = crate::test_support::bind_loopback_listener()
.await
.expect("listener should bind");
let addr = listener.local_addr().expect("local addr should resolve");
let app = Router::new().route(
"/rest/app-chat/conversations/new",
post(|body: Bytes| async move {
let body_json: serde_json::Value =
serde_json::from_slice(&body).expect("request body should be json");
if body_json.get("message").and_then(serde_json::Value::as_str)
!= Some("[user]: hello")
{
return (
axum::http::StatusCode::BAD_REQUEST,
Json(json!({
"error": {
"message": "expected grok app-chat message",
"body": body_json,
}
})),
);
}
(
axum::http::StatusCode::OK,
Json(json!({
"result": {
"response": {
"token": "pong",
"messageTag": "final"
}
}
})),
)
}),
);
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("test server should run");
});
let plan = ExecutionPlan {
request_id: "req-grok-runtime".into(),
candidate_id: Some("cand-grok".into()),
provider_name: Some("grok".into()),
provider_id: "provider-grok".into(),
endpoint_id: "endpoint-grok".into(),
key_id: "key-grok".into(),
method: "POST".into(),
url: format!("http://{addr}/rest/app-chat/conversations/new"),
headers: BTreeMap::from([
("content-type".into(), "application/json".into()),
(
aether_provider_transport::GROK_INTERNAL_HEADER.into(),
"1".into(),
),
]),
content_type: Some("application/json".into()),
content_encoding: None,
body: RequestBody::from_json(json!({
"model": "grok-4.20-0309-non-reasoning",
"messages": [{"role": "user", "content": "hello"}],
})),
stream: true,
client_api_format: "openai:chat".into(),
provider_api_format: "openai:chat".into(),
model_name: Some("grok-4.20-0309-non-reasoning".into()),
proxy: None,
transport_profile: None,
timeouts: Some(ExecutionTimeouts {
connect_ms: Some(5_000),
total_ms: Some(5_000),
..ExecutionTimeouts::default()
}),
};
let report_context = json!({"mapped_model": "grok-4.20-fast"});
let result = super::super::grok::maybe_execute_grok_sync(&plan, Some(&report_context))
.await
.expect("grok runtime plan should execute")
.expect("grok runtime should handle marked plan");
server.abort();
assert_eq!(result.status_code, http::StatusCode::OK.as_u16());
assert_eq!(
result
.body
.and_then(|body| body.json_body)
.and_then(|body| body["choices"][0]["message"]["content"]
.as_str()
.map(str::to_string)),
Some("pong".to_string())
);
}
#[tokio::test] #[tokio::test]
async fn execute_sync_plan_records_manual_proxy_success() { async fn execute_sync_plan_records_manual_proxy_success() {
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![ let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![

View File

@@ -38,7 +38,7 @@ pub(super) async fn build_admin_create_api_key_install_session_response(
Err(_) => { Err(_) => {
return Ok(build_admin_api_keys_bad_request_response( return Ok(build_admin_api_keys_bad_request_response(
"请求数据验证失败", "请求数据验证失败",
)) ));
} }
}; };

View File

@@ -16,6 +16,7 @@ use axum::{
Json, Json,
}; };
use serde_json::json; use serde_json::json;
use tracing::warn;
pub(super) async fn maybe_build_local_admin_payment_orders_response( pub(super) async fn maybe_build_local_admin_payment_orders_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
@@ -211,6 +212,19 @@ async fn build_admin_payment_credit_order_response(
.await? .await?
{ {
crate::AdminWalletMutationOutcome::Applied((order, credited)) => { crate::AdminWalletMutationOutcome::Applied((order, credited)) => {
if credited {
if let Err(err) = state
.app()
.apply_referral_rewards_for_payment_order_id(&order.id)
.await
{
warn!(
error = ?err,
order_id = %order.id,
"failed to apply referral rewards for admin-credited payment order"
);
}
}
Ok(attach_admin_audit_response( Ok(attach_admin_audit_response(
Json(json!({ Json(json!({
"order": build_admin_payment_order_payload(&order), "order": build_admin_payment_order_payload(&order),

View File

@@ -14,6 +14,7 @@ use axum::{
Json, Json,
}; };
use serde_json::json; use serde_json::json;
use tracing::warn;
pub(in super::super) async fn build_admin_wallet_complete_refund_response( pub(in super::super) async fn build_admin_wallet_complete_refund_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
@@ -86,6 +87,20 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
.await? .await?
{ {
crate::AdminWalletMutationOutcome::Applied(refund) => { crate::AdminWalletMutationOutcome::Applied(refund) => {
if let Some(order_id) = refund.payment_order_id.as_deref() {
if let Err(err) = state
.app()
.reverse_referral_rewards_for_order(order_id, refund.amount_usd)
.await
{
warn!(
error = ?err,
order_id = %order_id,
refund_id = %refund.id,
"failed to reverse referral rewards for completed refund"
);
}
}
let response = Json(json!({ let response = Json(json!({
"refund": build_admin_wallet_refund_payload(&wallet, &owner, &refund), "refund": build_admin_wallet_refund_payload(&wallet, &owner, &refund),
})) }))

View File

@@ -6,6 +6,8 @@ pub(super) mod features;
mod model; mod model;
pub(super) mod observability; pub(super) mod observability;
pub(super) mod provider; pub(super) mod provider;
mod referrals;
mod routing;
mod system; mod system;
mod users; mod users;

View File

@@ -6,6 +6,7 @@ use super::route_filters::{
}; };
use crate::constants::INTERNAL_GATEWAY_PATH_PREFIXES; use crate::constants::INTERNAL_GATEWAY_PATH_PREFIXES;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::build_admin_usage_counter_health_payload;
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::monitoring::{ use aether_admin::observability::monitoring::{
admin_monitoring_bad_request_response, admin_monitoring_user_behavior_user_id_from_path, admin_monitoring_bad_request_response, admin_monitoring_user_behavior_user_id_from_path,
@@ -189,6 +190,13 @@ pub(super) async fn build_admin_monitoring_system_status_response(
) )
.unwrap_or(usize::MAX); .unwrap_or(usize::MAX);
let tunnel = state.tunnel.stats(); let tunnel = state.tunnel.stats();
let usage_counter_snapshot = state
.data
.read_usage_counter_health()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let usage_counter =
build_admin_usage_counter_health_payload(&usage_counter_snapshot, now_unix_secs);
Ok(build_admin_monitoring_system_status_payload_response( Ok(build_admin_monitoring_system_status_payload_response(
now, now,
@@ -206,5 +214,6 @@ pub(super) async fn build_admin_monitoring_system_status_response(
tunnel.active_streams, tunnel.active_streams,
INTERNAL_GATEWAY_PATH_PREFIXES, INTERNAL_GATEWAY_PATH_PREFIXES,
recent_errors, recent_errors,
usage_counter,
)) ))
} }

View File

@@ -249,9 +249,10 @@ async fn admin_monitoring_resilience_status_returns_local_payload() {
let recommendations = payload["recommendations"] let recommendations = payload["recommendations"]
.as_array() .as_array()
.expect("recommendations should be array"); .expect("recommendations should be array");
assert!(recommendations.iter().any(|item| item assert!(recommendations.iter().any(|item| {
.as_str() item.as_str()
.is_some_and(|value| value.contains("prod-key")))); .is_some_and(|value| value.contains("prod-key"))
}));
assert!(payload["timestamp"].as_str().is_some()); assert!(payload["timestamp"].as_str().is_some());
} }

View File

@@ -1,7 +1,9 @@
use super::range::{build_comparison_range, parse_bounded_u32}; use super::range::{build_comparison_range, parse_bounded_u32};
use super::resolve_admin_usage_time_range; use super::resolve_admin_usage_time_range;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value}; use crate::handlers::admin::shared::{
build_admin_usage_counter_health_payload, query_param_optional_bool, query_param_value,
};
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::stats::{ use aether_admin::observability::stats::{
admin_stats_bad_request_response, admin_stats_comparison_empty_response, admin_stats_bad_request_response, admin_stats_comparison_empty_response,
@@ -21,6 +23,22 @@ use aether_data_contracts::repository::usage::{
}; };
use axum::{body::Body, http, response::Response}; use axum::{body::Body, http, response::Response};
async fn build_usage_counter_health_payload(
state: &AdminAppState<'_>,
) -> Result<serde_json::Value, GatewayError> {
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
let snapshot = state
.as_ref()
.data
.read_usage_counter_health()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
Ok(build_admin_usage_counter_health_payload(
&snapshot,
now_unix_secs,
))
}
fn usage_summary_to_admin_stats_aggregate( fn usage_summary_to_admin_stats_aggregate(
summary: &aether_data_contracts::repository::usage::StoredUsageAuditSummary, summary: &aether_data_contracts::repository::usage::StoredUsageAuditSummary,
) -> AdminStatsAggregate { ) -> AdminStatsAggregate {
@@ -211,13 +229,18 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
Ok(value) => u64::from(value.unwrap_or(10_000)), Ok(value) => u64::from(value.unwrap_or(10_000)),
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))), Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
}; };
let usage_counter = build_usage_counter_health_payload(state).await?;
if !state.has_usage_data_reader() { if !state.has_usage_data_reader() {
return Ok(Some(admin_stats_provider_performance_empty_response())); return Ok(Some(admin_stats_provider_performance_empty_response(
usage_counter,
)));
} }
let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds() let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds()
else { else {
return Ok(Some(admin_stats_provider_performance_empty_response())); return Ok(Some(admin_stats_provider_performance_empty_response(
usage_counter,
)));
}; };
let performance = state let performance = state
.summarize_usage_provider_performance(&UsageProviderPerformanceQuery { .summarize_usage_provider_performance(&UsageProviderPerformanceQuery {
@@ -240,6 +263,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
.await?; .await?;
return Ok(Some(build_admin_stats_provider_performance_response( return Ok(Some(build_admin_stats_provider_performance_response(
&performance, &performance,
usage_counter,
))); )));
} }

View File

@@ -5,12 +5,13 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::query_param_value; use crate::handlers::admin::shared::query_param_value;
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::usage::{ use aether_admin::observability::usage::{
admin_usage_bad_request_response, admin_usage_data_unavailable_response, admin_usage_bad_request_response, admin_usage_client_family,
admin_usage_has_fallback, admin_usage_is_failed, admin_usage_matches_search, admin_usage_data_unavailable_response, admin_usage_has_fallback, admin_usage_is_failed,
admin_usage_matches_username, admin_usage_parse_ids, admin_usage_parse_limit, admin_usage_matches_search, admin_usage_matches_username, admin_usage_parse_ids,
admin_usage_parse_offset, admin_usage_provider_key_name, admin_usage_record_json, admin_usage_parse_limit, admin_usage_parse_offset, admin_usage_provider_key_name,
build_admin_usage_active_requests_response, build_admin_usage_records_response, admin_usage_record_json, build_admin_usage_active_requests_response,
build_admin_usage_summary_stats_response_from_summary, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL, build_admin_usage_records_response, build_admin_usage_summary_stats_response_from_summary,
ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
}; };
use aether_data::repository::users::StoredUserSummary; use aether_data::repository::users::StoredUserSummary;
use aether_data_contracts::repository::{ use aether_data_contracts::repository::{
@@ -263,6 +264,19 @@ fn admin_usage_matches_attempt_status(
} }
} }
fn admin_usage_matches_client_family(
item: &StoredRequestUsageAudit,
client_family: Option<&str>,
) -> bool {
let Some(client_family) = client_family
.map(str::trim)
.filter(|value| !value.is_empty())
else {
return true;
};
admin_usage_client_family(item).is_some_and(|value| value.eq_ignore_ascii_case(client_family))
}
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
fn build_admin_usage_records_response_with_attempt_flags( fn build_admin_usage_records_response_with_attempt_flags(
items: &[StoredRequestUsageAudit], items: &[StoredRequestUsageAudit],
@@ -502,7 +516,9 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
.summarize_usage_audits(&UsageAuditSummaryQuery { .summarize_usage_audits(&UsageAuditSummaryQuery {
created_from_unix_secs, created_from_unix_secs,
created_until_unix_secs, created_until_unix_secs,
..Default::default() user_id: query_param_value(query, "user_id"),
provider_name: query_param_value(query, "provider"),
model: query_param_value(query, "model"),
}) })
.await?; .await?;
return Ok(Some(build_admin_usage_summary_stats_response_from_summary( return Ok(Some(build_admin_usage_summary_stats_response_from_summary(
@@ -598,6 +614,7 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
admin_usage_attempt_status_filter(query_param_value(query, "status").as_deref()); admin_usage_attempt_status_filter(query_param_value(query, "status").as_deref());
let search = query_param_value(query, "search"); let search = query_param_value(query, "search");
let username_filter = query_param_value(query, "username"); let username_filter = query_param_value(query, "username");
let client_family_filter = query_param_value(query, "client_family");
let limit = match admin_usage_parse_limit(query) { let limit = match admin_usage_parse_limit(query) {
Ok(value) => value, Ok(value) => value,
Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))), Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))),
@@ -632,7 +649,12 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
let active_username_filter = username_filter let active_username_filter = username_filter
.as_deref() .as_deref()
.filter(|value| !value.trim().is_empty()); .filter(|value| !value.trim().is_empty());
let (usage, total) = if let Some(attempt_status) = attempt_status_filter { let active_client_family_filter = client_family_filter
.as_deref()
.filter(|value| !value.trim().is_empty());
let (usage, total) = if attempt_status_filter.is_some()
|| active_client_family_filter.is_some()
{
let mut usage = state.list_usage_audits(&base_query).await?; let mut usage = state.list_usage_audits(&base_query).await?;
let user_ids: Vec<String> = usage let user_ids: Vec<String> = usage
.iter() .iter()
@@ -662,12 +684,14 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
active_username_filter, active_username_filter,
&users_by_id, &users_by_id,
state.has_auth_user_data_reader(), state.has_auth_user_data_reader(),
) && admin_usage_matches_attempt_status( ) && attempt_status_filter.is_none_or(|attempt_status| {
item, admin_usage_matches_attempt_status(
attempt_status, item,
&attempt_flags_by_usage_id, attempt_status,
request_candidate_reader_available, &attempt_flags_by_usage_id,
) request_candidate_reader_available,
)
}) && admin_usage_matches_client_family(item, active_client_family_filter)
}); });
sort_usage_newest_first(&mut usage); sort_usage_newest_first(&mut usage);
let total = usage.len(); let total = usage.len();

View File

@@ -13,12 +13,18 @@ pub(super) fn key_api_formats_without_entry(
} }
pub(super) fn endpoint_key_counts_by_format( pub(super) fn endpoint_key_counts_by_format(
provider_type: &str,
endpoints: &[StoredProviderCatalogEndpoint],
keys: &[StoredProviderCatalogKey], keys: &[StoredProviderCatalogKey],
) -> ( ) -> (
std::collections::BTreeMap<String, usize>, std::collections::BTreeMap<String, usize>,
std::collections::BTreeMap<String, usize>, std::collections::BTreeMap<String, usize>,
) { ) {
admin_provider_endpoints_pure::endpoint_key_counts_by_format(keys) admin_provider_endpoints_pure::endpoint_key_counts_by_format(provider_type, endpoints, keys)
}
pub(super) fn normalize_endpoint_api_format(api_format: &str) -> String {
admin_provider_endpoints_pure::normalize_endpoint_api_format(api_format)
} }
pub(super) fn build_admin_provider_endpoint_response( pub(super) fn build_admin_provider_endpoint_response(

View File

@@ -4,7 +4,10 @@ use aether_data_contracts::repository::provider_catalog::{
}; };
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use super::payloads::{build_admin_provider_endpoint_response, endpoint_key_counts_by_format}; use super::payloads::{
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
normalize_endpoint_api_format,
};
pub(crate) async fn build_admin_provider_endpoints_payload( pub(crate) async fn build_admin_provider_endpoints_payload(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
@@ -38,7 +41,8 @@ pub(crate) async fn build_admin_provider_endpoints_payload(
.await .await
.ok() .ok()
.unwrap_or_default(); .unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys); let (total_keys_by_format, active_keys_by_format) =
endpoint_key_counts_by_format(&provider.provider_type, &endpoints, &keys);
let now_unix_secs = SystemTime::now() let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH) .duration_since(UNIX_EPOCH)
.ok() .ok()
@@ -51,15 +55,16 @@ pub(crate) async fn build_admin_provider_endpoints_payload(
.skip(skip) .skip(skip)
.take(limit) .take(limit)
.map(|endpoint| { .map(|endpoint| {
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
build_admin_provider_endpoint_response( build_admin_provider_endpoint_response(
&endpoint, &endpoint,
&provider.name, &provider.name,
total_keys_by_format total_keys_by_format
.get(endpoint.api_format.as_str()) .get(endpoint_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
active_keys_by_format active_keys_by_format
.get(endpoint.api_format.as_str()) .get(endpoint_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
now_unix_secs, now_unix_secs,
@@ -92,22 +97,27 @@ pub(crate) async fn build_admin_endpoint_payload(
.await .await
.ok() .ok()
.unwrap_or_default(); .unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys); let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(
&provider.provider_type,
std::slice::from_ref(&endpoint),
&keys,
);
let now_unix_secs = SystemTime::now() let now_unix_secs = SystemTime::now()
.duration_since(UNIX_EPOCH) .duration_since(UNIX_EPOCH)
.ok() .ok()
.map(|duration| duration.as_secs()) .map(|duration| duration.as_secs())
.unwrap_or(0); .unwrap_or(0);
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
Some(build_admin_provider_endpoint_response( Some(build_admin_provider_endpoint_response(
&endpoint, &endpoint,
&provider.name, &provider.name,
total_keys_by_format total_keys_by_format
.get(endpoint.api_format.as_str()) .get(endpoint_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
active_keys_by_format active_keys_by_format
.get(endpoint.api_format.as_str()) .get(endpoint_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
now_unix_secs, now_unix_secs,

View File

@@ -1,7 +1,7 @@
use super::extractors::admin_endpoint_id; use super::extractors::admin_endpoint_id;
use super::payloads::{ use super::payloads::{
build_admin_provider_endpoint_response, endpoint_key_counts_by_format, build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
AdminProviderEndpointUpdatePatch, normalize_endpoint_api_format, AdminProviderEndpointUpdatePatch,
}; };
use super::support::build_admin_endpoints_data_unavailable_response; use super::support::build_admin_endpoints_data_unavailable_response;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
@@ -147,18 +147,23 @@ pub(super) async fn maybe_handle(
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id)) .list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await .await
.unwrap_or_default(); .unwrap_or_default();
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys); let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(
&provider.provider_type,
std::slice::from_ref(&updated),
&keys,
);
let updated_api_format = normalize_endpoint_api_format(&updated.api_format);
Ok(Some( Ok(Some(
Json(build_admin_provider_endpoint_response( Json(build_admin_provider_endpoint_response(
&updated, &updated,
&provider.name, &provider.name,
total_keys_by_format total_keys_by_format
.get(updated.api_format.as_str()) .get(updated_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
active_keys_by_format active_keys_by_format
.get(updated.api_format.as_str()) .get(updated_api_format.as_str())
.copied() .copied()
.unwrap_or(0), .unwrap_or(0),
now_unix_secs, now_unix_secs,

View File

@@ -32,7 +32,7 @@ pub(super) async fn maybe_handle(
.into_response(), .into_response(),
)); ));
}; };
let Some(_provider) = state let Some(provider) = state
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id)) .read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
.await? .await?
.into_iter() .into_iter()
@@ -129,7 +129,9 @@ pub(super) async fn maybe_handle(
Json(serde_json::Value::Array( Json(serde_json::Value::Array(
created created
.iter() .iter()
.map(|model| build_admin_provider_model_response(model, now_unix_secs)) .map(|model| {
build_admin_provider_model_response(&provider, model, now_unix_secs)
})
.collect(), .collect(),
)) ))
.into_response(), .into_response(),

Some files were not shown because too many files have changed in this diff Show More