From 843376f70cdd5447b681187c21371fb332e27893 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=A2=A8=E8=8F=8A?= <277378677+maiphucgiang@users.noreply.github.com> Date: Wed, 16 Sep 2026 23:46:35 +0800 Subject: [PATCH] Add opt-in upstream pooling and account capacity limits Own bounded cookie-free clients in the ASGI lifespan and release account-scoped leases on completion, failover and cancellation. Preserve strict routes, free-first admission and existing replay rules. Expose compatible defaults through CLI, environment and WebUI settings. Validate all 51 backend test scripts, real loopback TCP reuse, and isolated native deployment with real client disconnects. Keep version 1.2.6. --- .env.example | 4 + app/inference_resources.py | 156 +++++++++++ app/settings.py | 4 + app/upstream_io.py | 16 +- converter.py | 66 ++++- docker-compose.yml | 2 + docs/advanced.md | 11 + docs/advanced.zh-CN.md | 11 + tests/test_environment_config.py | 24 +- tests/test_inference_resources.py | 417 ++++++++++++++++++++++++++++++ 10 files changed, 698 insertions(+), 13 deletions(-) create mode 100644 app/inference_resources.py create mode 100644 tests/test_inference_resources.py diff --git a/.env.example b/.env.example index af30ada..4f89cd9 100644 --- a/.env.example +++ b/.env.example @@ -35,6 +35,10 @@ CODEBUDDY2API_MAX_INBOUND_BYTES=67108864 # Aggregated output bytes and concurrent inference requests; zero disables each limit. CODEBUDDY2API_MAX_COLLECT_BYTES=8388608 CODEBUDDY2API_MAX_CONCURRENT=64 +# Optional bounded upstream connection reuse; restart to apply, disabled by default. +# CODEBUDDY2API_UPSTREAM_KEEPALIVE=false +# Per-account in-flight limit; zero preserves unlimited account capacity. +# CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT=0 # SQLite audit defaults to the data directory; this optional path enables a separate text log. CODEBUDDY2API_LOG= diff --git a/app/inference_resources.py b/app/inference_resources.py new file mode 100644 index 0000000..40d780a --- /dev/null +++ b/app/inference_resources.py @@ -0,0 +1,156 @@ +"""Own bounded upstream clients and account capacity for each inference request.""" +from __future__ import annotations + +import asyncio +from contextlib import asynccontextmanager +from contextvars import ContextVar +from http.cookiejar import CookieJar, DefaultCookiePolicy +import threading + +import httpx + +from app.site_routing import PROFILE_ENDPOINTS + + +class _RejectCookies(DefaultCookiePolicy): + def set_ok(self, cookie, request): + return False + + def return_ok(self, cookie, request): + return False + + +class UpstreamClients: + """Keep at most one cookie-free HTTP/1.1 client per trusted origin and event loop.""" + + def __init__(self): + self._clients = {} + self._loop = asyncio.get_running_loop() + self._closed = False + self._origins = {self._origin(url) for url in PROFILE_ENDPOINTS.values()} + + @staticmethod + def _origin(url): + value = httpx.URL(url) + return value.scheme, value.host, value.port + + def get(self, url): + origin = self._origin(url) + if self._closed or asyncio.get_running_loop() is not self._loop or origin not in self._origins: + return None + if origin not in self._clients: + self._clients[origin] = httpx.AsyncClient( + cookies=CookieJar(policy=_RejectCookies()), http2=False, + limits=httpx.Limits(max_connections=64, max_keepalive_connections=16, keepalive_expiry=30)) + return self._clients[origin] + + async def aclose(self): + self._closed = True + clients, self._clients = list(self._clients.values()), {} + results = await asyncio.gather(*(client.aclose() for client in clients), return_exceptions=True) + errors = [result for result in results if isinstance(result, BaseException)] + if errors: + raise BaseExceptionGroup("Upstream client shutdown failed", errors) + + +@asynccontextmanager +async def inference_lifespan(app): + clients = UpstreamClients() + try: + yield {"upstream_clients": clients} + finally: + await clients.aclose() + + +class CredentialLease(tuple): + """Retain the existing (manager, generation) lease shape with idempotent capacity release.""" + + def __new__(cls, manager, generation, release): + lease = super().__new__(cls, (manager, generation)) + lease._release = release + lease._lock = threading.Lock() + return lease + + def release(self): + with self._lock: + release, self._release = self._release, None + if release is not None: + release() + + +def release_credential(credential): + if isinstance(credential, CredentialLease): + credential.release() + + +class AccountCapacity: + """Count account leases independently of credential I/O locks.""" + + def __init__(self): + self._lock = threading.Lock() + self._counts = {} + + def count(self, identity): + with self._lock: + return self._counts.get(identity, 0) + + def acquire(self, identity, limit, manager, generation): + with self._lock: + count = self._counts.get(identity, 0) + if limit and count >= limit: + return None + self._counts[identity] = count + 1 + return CredentialLease(manager, generation, lambda: self._release(identity)) + + def _release(self, identity): + with self._lock: + count = self._counts[identity] - 1 + if count: + self._counts[identity] = count + else: + del self._counts[identity] + + +class RequestResources: + """Release even leases acquired by a worker after its request has already closed.""" + + def __init__(self, clients=None): + self.clients = clients + self._lock = threading.Lock() + self._leases = [] + self._closed = False + + def add(self, lease): + with self._lock: + if not self._closed: + self._leases.append(lease) + return + release_credential(lease) + raise asyncio.CancelledError() + + def close(self): + with self._lock: + self._closed = True + leases, self._leases = self._leases, [] + for lease in leases: + release_credential(lease) + + +request_resources = ContextVar("inference_resources", default=None) + + +class InferenceResourcesMiddleware: + def __init__(self, app): + self.app = app + + async def __call__(self, scope, receive, send): + if (scope["type"] != "http" or scope.get("method") != "POST" or + scope.get("path") not in ("/v1/chat/completions", "/v1/responses", "/v1/messages")): + return await self.app(scope, receive, send) + resources = RequestResources(scope.get("state", {}).get("upstream_clients")) + token = request_resources.set(resources) + try: + await self.app(scope, receive, send) + finally: + resources.close() + request_resources.reset(token) diff --git a/app/settings.py b/app/settings.py index c180bfb..9410ec2 100644 --- a/app/settings.py +++ b/app/settings.py @@ -43,6 +43,10 @@ def _item(default, type_, label, *, mode="hot", env=None, minimum=None, maximum= minimum=0, maximum=10), "retry_write_timeout": _item(False, "boolean", "写超时参与重放", env="CODEBUDDY2API_RETRY_WRITE_TIMEOUT"), + "upstream_keepalive": _item(False, "boolean", "上游连接复用", mode="restart", + env="CODEBUDDY2API_UPSTREAM_KEEPALIVE"), + "max_inflight_per_account": _item(0, "integer", "单账号在途上限(0 不限制)", + env="CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT", minimum=0, maximum=10000), "audit_max_bytes": _item(256 * 1024 * 1024, "integer", "审计明细预算", minimum=1024**2, maximum=1024**4), "audit_retention_days": _item(30, "integer", "审计明细保留天数", minimum=1, maximum=36500), "audit_diagnostic_bytes": _item(8192, "integer", "失败诊断最大字节", minimum=0, maximum=8192), diff --git a/app/upstream_io.py b/app/upstream_io.py index 8ee12f9..226f171 100644 --- a/app/upstream_io.py +++ b/app/upstream_io.py @@ -211,9 +211,19 @@ async def read_bounded_error(response, limit: int = ERROR_BODY_LIMIT) -> bytes: WRITE_TIMEOUT = (httpx.WriteTimeout,) +@asynccontextmanager +async def _attempt_client(url, timeout, clients): + client = clients.get(url) if clients is not None else None + if client is not None: + yield client + else: + async with httpx.AsyncClient(timeout=timeout) as client: + yield client + + @asynccontextmanager async def open_backend_stream(url, headers, body, *, read_timeout=300, on_retry=None, - retry_write_timeout=False): + retry_write_timeout=False, clients=None): """Retry connection failures once on a fresh client; write timeouts require explicit opt-in. Never replay after the upstream response opens. """ @@ -222,8 +232,8 @@ async def open_backend_stream(url, headers, body, *, read_timeout=300, on_retry= for attempt in range(2): opened = False try: - async with httpx.AsyncClient(timeout=timeout) as client: - async with client.stream("POST", url, headers=headers, json=body) as response: + async with _attempt_client(url, timeout, clients if attempt == 0 else None) as client: + async with client.stream("POST", url, headers=headers, json=body, timeout=timeout) as response: opened = True yield response return diff --git a/converter.py b/converter.py index 3750dc1..8e77311 100644 --- a/converter.py +++ b/converter.py @@ -59,6 +59,8 @@ def desensitize_body(body, roles=("system",), desensitize_harness_user=False, credential_file_lock) from app.upstream_io import (ChatSSEAccumulator, UpstreamHTTPError, UpstreamResponseError, open_backend_stream, parse_retry_after, read_bounded_error) +from app.inference_resources import (AccountCapacity, InferenceResourcesMiddleware, inference_lifespan, + request_resources, release_credential) from app.inference_auth import require_api_key from app.content_filter import ContentFilterDetector, is_filter_error from app.request_limits import ImageLimitError, apply_image_policy @@ -464,6 +466,7 @@ def __init__(self, paths: list[Path] | None = None, scan: bool = False, self._blocks = ModelBlocks(blocks_path, ttl_s=MODEL_SITE_BLOCK_S, max_ttl_s=MODEL_SITE_BLOCK_MAX_S) self._rr = {None: 0, "cn": 0, "intl": 0} self._ledger = None # Prefer credits expiring sooner. + self._capacity = AccountCapacity() self._scan = scan # Rescan credentials before selection. self._ignored_duplicates: set[str] = set() self._sync_pending: set[str] = set() @@ -793,8 +796,19 @@ def _candidates(self, model: str | None, *, region=None, tried=()) -> list[dict] healthy.sort(key=lambda entry: (not self._model_free(entry, model), *self._expiry_rank(entry))) return healthy + @staticmethod + def _capacity_error(): + return HTTPException(status_code=503, headers={"Retry-After": "3"}, detail={"error": { + "message": "符合当前路由和免费优先策略的账号在途名额已满,请稍后重试", + "type": "service_unavailable", "code": "credential_concurrency_limit"}}) + + @staticmethod + def _capacity_key(entry): + return entry.get("account_key") or entry["id"] + + def pick(self, skey: str | None, model: str | None = None, *, region=None, - tried=()) -> CredentialManager | None: + tried=(), with_capacity=False) -> CredentialManager | None: """Select a healthy sticky or round-robin credential, preferring eligible zero-rate accounts.""" self._rescan() # Reload and prune acquire their own locks. with self._lock: @@ -804,6 +818,13 @@ def pick(self, skey: str | None, model: str | None = None, *, region=None, if skey: self._sticky.pop(skey, None) return None + limit = CONFIG.get("max_inflight_per_account", 0) + if with_capacity and limit: + free = self._model_free(candidates[0], model) + candidates = [entry for entry in candidates if self._model_free(entry, model) == free + and self._capacity.count(self._capacity_key(entry)) < limit] + if not candidates: + raise self._capacity_error() best = candidates[0] free = self._model_free(best, model) top = [e for e in candidates if self._model_free(e, model) == free @@ -822,10 +843,11 @@ def pick(self, skey: str | None, model: str | None = None, *, region=None, return e["cm"] def headers_for(self, skey: str | None, model: str | None = None, *, region=None, - with_generation=False, tried=()): - """Recheck the credential generation and site before sending.""" + with_generation=False, tried=(), with_capacity=False): + """Recheck identity and atomically reserve account capacity before sending.""" + capacity_race = False for _ in range(max(1, len(self._entries))): - cm = self.pick(skey, model, region=region, tried=tried) + cm = self.pick(skey, model, region=region, tried=tried, with_capacity=with_capacity) if cm is None: return None reason = None @@ -844,7 +866,16 @@ def headers_for(self, skey: str | None, model: str | None = None, *, region=None entry = next((entry for entry in self._entries if entry["cm"] is cm), None) if (entry is not None and cm._generation == generation and self._healthy(entry) and self._eligible(entry, model, region=region, profile=profile) and self._model_healthy(entry, model)): + if with_capacity: + lease = self._capacity.acquire(self._capacity_key(entry), + CONFIG.get("max_inflight_per_account", 0), cm, generation) + if lease is None: + capacity_race = True + continue + return lease, headers return ((cm, generation) if with_generation else cm), headers + if capacity_race: + raise self._capacity_error() return None @staticmethod @@ -1041,6 +1072,8 @@ def snapshot(self) -> list[dict]: out = [] for e in self._entries: s: dict = {"auth_file": e["id"], "healthy": self._healthy(e), + "in_flight": self._capacity.count(self._capacity_key(e)), + "max_in_flight": CONFIG.get("max_inflight_per_account", 0), "model_cooldowns": {m: time.strftime("%m-%d %H:%M:%S", time.localtime(u)) for (cid, m), u in self._model_fail.items() if cid == e["id"] and u > now}, @@ -1401,7 +1434,8 @@ def _housekeeper_loop(pool: CredentialPool, ledger) -> None: # FastAPI application # --------------------------------------------------------------------------- -app = FastAPI(title="codebuddy2api", version=APP_VERSION) +app = FastAPI(title="codebuddy2api", version=APP_VERSION, lifespan=inference_lifespan) +app.add_middleware(InferenceResourcesMiddleware) # Anthropic error types: https://platform.claude.com/docs/en/api/errors _ANTHROPIC_ERROR_TYPES = { @@ -1450,6 +1484,7 @@ async def _protocol_http_exception(request: Request, exc: HTTPException): "max_request_bytes": 32 * 1024 * 1024, "log_body_limit": 65536, "max_inbound_bytes": 64 * 1024 * 1024, "max_collect_bytes": 8 * 1024 * 1024, "max_concurrent": 64, + "upstream_keepalive": False, "max_inflight_per_account": 0, "failover_max": 0, # Credential failovers allowed before the first response byte "retry_write_timeout": False, # Opt-in replay after incomplete writes "usage_daily": None, # Usage aggregated by date and model @@ -1541,7 +1576,9 @@ def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=()) skey = model_policy.sticky_scope(CONFIG, skey, model) pool = CONFIG.get("cred_pool") if pool is not None: - picked = pool.headers_for(skey, model, region=region, with_generation=True, tried=tried) + resources = request_resources.get() + picked = pool.headers_for(skey, model, region=region, with_generation=True, tried=tried, + with_capacity=resources is not None) if picked is None: until = pool.model_cooldown_until(model, region=region) if until: @@ -1564,6 +1601,8 @@ def _cred_for(payload: dict, model: str | None = None, *, region=None, tried=()) detail={"error": {"message": "无可用凭证(未登录、目录/额度未就绪或全部熔断)", "type": "auth_error"}}) cm, headers = picked + if resources is not None: + resources.add(cm) else: cm = CONFIG["cred"] if cm is None or cm in {_cred_manager(item) for item in tried}: @@ -2589,8 +2628,11 @@ def retry(error): _log(f"[{rid}] {'写超时重放' if timeout_on_write else '建连失败'},重试 1/1 | {model_name}" f" | {_network_error_text(error)}{_replay_cost_note(error)}") try: + resources = request_resources.get() + clients = resources.clients if resources is not None and CONFIG.get("upstream_keepalive") else None async with open_backend_stream(url, headers, body, read_timeout=timeout, on_retry=retry, - retry_write_timeout=bool(CONFIG.get("retry_write_timeout"))) as response: + retry_write_timeout=bool(CONFIG.get("retry_write_timeout")), + clients=clients) as response: opened = True observe_attempt("upstream_http", status_code=response.status_code, duration_ms=(time.monotonic() - started) * 1000) @@ -2928,6 +2970,7 @@ async def _stream_plan(payload, canonical, model_name, rid, t0, make, routed, cr first = await _preflight_stream(stream, model_name, t0, rid) except _StreamFailure as failure: await _close_stream(stream) # Release the failed upstream connection. + release_credential(cred) recovered = observe_failure_seq() # Recover only this failure sequence. tried.append(cred) limit = _failover_limit() @@ -2976,6 +3019,7 @@ async def _routed_fetch(payload, canonical, model_name, rid, t0, fetch, routed, observe_recovery(recovered) return collected except (httpx.HTTPError, UpstreamResponseError) as error: + release_credential(cred) status, raw = _upstream_failure(error, model_name, t0, rid) recovered = observe_failure_seq() tried.append(cred) @@ -3390,6 +3434,12 @@ def main(): ap.add_argument("--max-concurrent", type=_nonnegative_int, metavar="N", default=os.environ.get("CODEBUDDY2API_MAX_CONCURRENT", "64"), help="推理端点并发上限(超出立即 503),默认 64;0 不限制") + ap.add_argument("--max-inflight-per-account", type=_nonnegative_int, metavar="N", + default=os.environ.get("CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT", "0"), + help="单账号在途上限,默认 0(不限制);满载返回 503,不借容量切换到收费账号") + ap.add_argument("--upstream-keepalive", type=_boolean_arg, nargs="?", const=True, + default=os.environ.get("CODEBUDDY2API_UPSTREAM_KEEPALIVE", "false"), + help="按上游入口复用有界连接池,默认 false;重启生效,不改变超时或重放规则") ap.add_argument("--log-body-limit", type=_nonnegative_int, metavar="BYTES", default=os.environ.get("CODEBUDDY2API_LOG_BODY_LIMIT", "65536"), help="每条正文日志的预览字节上限,默认 64 KiB;0 只记录摘要") @@ -3419,7 +3469,7 @@ def main(): for key in ("max_images", "image_policy", "max_request_bytes", "log_body_limit", "tool_call_max_retry", "max_inbound_bytes", "max_collect_bytes", "max_concurrent", - "failover_max", "retry_write_timeout"): + "failover_max", "retry_write_timeout", "upstream_keepalive", "max_inflight_per_account"): CONFIG[key] = getattr(args, key) CONFIG["api_key"] = args.api_key CONFIG["desensitize"] = args.desensitize diff --git a/docker-compose.yml b/docker-compose.yml index 916eaf3..7913d8e 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -27,6 +27,8 @@ services: CODEBUDDY2API_MAX_INBOUND_BYTES: ${CODEBUDDY2API_MAX_INBOUND_BYTES:-67108864} CODEBUDDY2API_MAX_COLLECT_BYTES: ${CODEBUDDY2API_MAX_COLLECT_BYTES:-8388608} CODEBUDDY2API_MAX_CONCURRENT: ${CODEBUDDY2API_MAX_CONCURRENT:-64} + CODEBUDDY2API_UPSTREAM_KEEPALIVE: + CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT: CODEBUDDY2API_TOOL_CALL_MAX_RETRY: ${CODEBUDDY2API_TOOL_CALL_MAX_RETRY:-3} CODEBUDDY2API_FAILOVER_MAX: CODEBUDDY2API_RETRY_WRITE_TIMEOUT: diff --git a/docs/advanced.md b/docs/advanced.md index c6dfb7c..f47d661 100644 --- a/docs/advanced.md +++ b/docs/advanced.md @@ -32,6 +32,8 @@ Compose explicitly passes some environment variables and CLI flags, so deleting | `--max-inbound-bytes` | `67108864` | Raw body limit for generation and token-count POSTs, before parsing (chunked included); other routes are not buffered; 413 beyond it | | `--max-collect-bytes` | `8388608` | Total collection budget for aggregated output (content + reasoning + tool arguments); `response_too_large` beyond it; `0` disables | | `--max-concurrent` | `64` | Concurrency limit for the three generation endpoints only; excess requests get 503 with Retry-After; token counting is unaffected; `0` disables | +| `--max-inflight-per-account` | `0` | Per-process, per-account in-flight client inference limit; `0` disables, full accounts return 503 | +| `--upstream-keepalive [true/false]` | `false` | Bounded connection reuse isolated by official origin; requires restart | | `--failover-max` | `0` | Extra credentials tried when a request fails before the first response byte reaches the client; `0` keeps the upstream behaviour of surfacing the failure directly | | `--retry-write-timeout` | `false` | Opt a request-body write timeout into replay (fresh connection and `--failover-max`), accepting that bytes already sent may have been processed | | `--max-request-bytes` | `33554432` | Positive byte limit for the processed upstream JSON | @@ -61,6 +63,15 @@ Off by default, preserving the existing policy: desensitization strips tool desc Use a source/image build and Compose configuration containing this feature; recreate containers after changing their environment. Retained descriptions may increase input tokens and content-filter rejections; compatibility across accounts/models is not guaranteed. Set `false` to restore the previous policy. This option does not restore other schema fields or deep nodes removed by existing Responses projection, nor relax the request-size budget. +### Connection reuse and account capacity + +Both settings are available in the WebUI; their environment variables are `CODEBUDDY2API_UPSTREAM_KEEPALIVE` and `CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT`. Unset variables leave Compose settings unlocked. Connection reuse defaults to off; when enabled, each official origin permits 64 connections with 16 idle connections and a 30-second keepalive expiry. Authentication is request-scoped, upstream cookies are not stored, and shutdown closes the pools. Proxy environment, timeouts and replay rules are unchanged; disable and restart to restore fresh connections. + +The account limit defaults to `0`. Positive limits skip full accounts within existing routing and free-first rules; a full free tier never spills into paid accounts. No capacity returns `503 / credential_concurrency_limit` with `Retry-After: 3`, without queueing or penalizing the account. Completion, cancellation and failed-account rotation release capacity. The credentials API exposes `in_flight` and `max_in_flight`. Only the three client generation endpoints count; limits are per process, not shared between instances. Setting `0` restores unlimited account capacity without interrupting active requests. + +Before downgrading the source, remove the new CLI arguments and restore a control-store backup without these keys; disabling the features does not remove persisted settings. + + ## APIs and authentication | Client endpoint | Description | diff --git a/docs/advanced.zh-CN.md b/docs/advanced.zh-CN.md index db5a831..4be22de 100644 --- a/docs/advanced.zh-CN.md +++ b/docs/advanced.zh-CN.md @@ -32,6 +32,8 @@ Compose 会显式传入部分环境变量及 CLI 参数,删除 `.env` 中的 | `--max-inbound-bytes` | `67108864` | 生成及 token 估算 POST 的解析前原始字节上限(含 chunked),超限 413;其他路由不缓冲请求体 | | `--max-collect-bytes` | `8388608` | 聚合路径输出收集总字节上限(正文+思考+工具参数),超限返回 `response_too_large`;`0` 不限制 | | `--max-concurrent` | `64` | 仅限制三个生成端点;占满立即 503(含 Retry-After),不限制 token 估算;`0` 不限制 | +| `--max-inflight-per-account` | `0` | 每进程、每账号的客户端推理在途上限;`0` 不限制,满载立即 503 | +| `--upstream-keepalive [true/false]` | `false` | 启用按官方入口隔离的有界连接复用;重启生效 | | `--failover-max` | `0` | 请求在「一个字节都还没发给下游」之前失败时,最多再换几个凭证就地重放;`0` 表示如实把失败回给下游 | | `--retry-write-timeout` | `false` | 让「写请求体超时」也参与重放(换新连接与 `--failover-max` 换凭证),代价是已发出的那半截正文可能已被上游处理 | | `--max-request-bytes` | `33554432` | 处理后的上游 JSON 字节上限,须为正整数 | @@ -61,6 +63,15 @@ Compose 会显式传入部分环境变量及 CLI 参数,删除 `.env` 中的 需使用包含此功能的源码/镜像和 Compose 配置;修改容器环境后重新创建容器。保留描述可能增加输入 token 和审核拦截风险,不保证所有账号/模型都同样兼容;设为 `false` 可恢复旧策略。此开关不恢复 Responses 原有投影裁掉的其他 schema 字段或深层节点,也不放宽请求体预算。 +### 连接复用与账号容量 + +WebUI 系统设置可配置这两项;环境变量为 `CODEBUDDY2API_UPSTREAM_KEEPALIVE` 和 `CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT`,未设置时 Compose 不锁定 WebUI。连接复用默认关闭,启用后每个官方入口最多 64 条连接、保留 16 条空闲连接,空闲复用期限 30 秒;认证头逐请求设置,不保存上游 Cookie,关闭服务时释放连接池。原有代理环境、超时和重放规则不变;关闭并重启恢复逐请求连接。 + +账号上限默认 `0`;设为正数后,仅在原路由及免费优先范围内避开满载账号,不因免费账号满载而转向收费账号。没有名额时返回 `503 / credential_concurrency_limit` 和 `Retry-After: 3`,不排队、不熔断;结束、断连和失败换号均释放名额。管理凭据 API 提供 `in_flight`、`max_in_flight`;限制只涵盖三个客户端生成接口,每进程独立,多个实例不共享计数。设回 `0` 即恢复原容量策略,不中断已开始的请求。 + +源码降级前还需移除新增启动参数,并恢复不含这两个配置键的控制库备份;仅关闭开关不会删除持久化配置。 + + ## API 与鉴权 | 客户端接口 | 说明 | diff --git a/tests/test_environment_config.py b/tests/test_environment_config.py index f970b05..5fbb2a6 100644 --- a/tests/test_environment_config.py +++ b/tests/test_environment_config.py @@ -94,12 +94,30 @@ def test_environment_public_binding_without_key_is_rejected(self): def test_existing_runtime_limit_and_retry_variables_reach_effective_config(self): values = {'MAX_INBOUND_BYTES': '4096', 'MAX_COLLECT_BYTES': '0', 'MAX_CONCURRENT': '2', - 'TOOL_CALL_MAX_RETRY': '1', 'FAILOVER_MAX': '1', 'RETRY_WRITE_TIMEOUT': 'true'} + 'TOOL_CALL_MAX_RETRY': '1', 'FAILOVER_MAX': '1', 'RETRY_WRITE_TIMEOUT': 'true', + 'UPSTREAM_KEEPALIVE': 'true', 'MAX_INFLIGHT_PER_ACCOUNT': '2'} _, _, config = self.start({'CODEBUDDY2API_' + name: value for name, value in values.items()}) for name, value in values.items(): expected = value == 'true' if value in ('true', 'false') else int(value) self.assertEqual(config[name.lower()], expected) + def test_pooling_cli_wins_over_environment_and_saved_settings(self): + _, items, config = self.start( + {'CODEBUDDY2API_UPSTREAM_KEEPALIVE': 'true', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT': '3'}, + cli=('--upstream-keepalive=false', '--max-inflight-per-account', '1'), + saved={'upstream_keepalive': True, 'max_inflight_per_account': 2}) + self.assertFalse(config['upstream_keepalive']) + self.assertEqual(config['max_inflight_per_account'], 1) + self.assertTrue(all(items[key]['source'] == 'cli' and items[key]['locked'] + for key in ('upstream_keepalive', 'max_inflight_per_account'))) + + def test_invalid_pooling_environment_fails_before_startup(self): + for env in ({'CODEBUDDY2API_UPSTREAM_KEEPALIVE': 'invalid'}, + {'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT': '-1'}): + with self.subTest(env=env), self.assertRaises(SystemExit): + self.start(env) + + def test_example_covers_all_active_runtime_environment_names(self): example = (ROOT / '.env.example').read_text() documented = set(re.findall(r'(?m)^(?:# )?(CODEBUDDY[A-Z0-9_]+)=', example)) @@ -132,6 +150,7 @@ def test_compose_forwards_dotenv_limits_and_retries_without_changing_internal_bi 'CODEBUDDY2API_MAX_INBOUND_BYTES': '4096', 'CODEBUDDY2API_MAX_COLLECT_BYTES': '0', 'CODEBUDDY2API_MAX_CONCURRENT': '2', 'CODEBUDDY2API_TOOL_CALL_MAX_RETRY': '1', 'CODEBUDDY2API_FAILOVER_MAX': '1', 'CODEBUDDY2API_RETRY_WRITE_TIMEOUT': 'true', + 'CODEBUDDY2API_UPSTREAM_KEEPALIVE': 'true', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT': '2', 'CODEBUDDY2API_KEEP_TOOL_METADATA': 'false', 'CODEBUDDY_IMPORT_DIR': '/data/auth/incoming'} service = self.compose(values) port = service['ports'][0] @@ -144,7 +163,8 @@ def test_compose_forwards_dotenv_limits_and_retries_without_changing_internal_bi def test_compose_unset_optional_settings_do_not_override_webui(self): service = self.compose({}) - for name in ('CODEBUDDY2API_KEEP_TOOL_METADATA', 'CODEBUDDY2API_FAILOVER_MAX', 'CODEBUDDY2API_RETRY_WRITE_TIMEOUT'): + for name in ('CODEBUDDY2API_KEEP_TOOL_METADATA', 'CODEBUDDY2API_FAILOVER_MAX', 'CODEBUDDY2API_RETRY_WRITE_TIMEOUT', + 'CODEBUDDY2API_UPSTREAM_KEEPALIVE', 'CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT'): self.assertIsNone(service['environment'].get(name)) self.assertEqual(service['ports'][0]['host_ip'], '127.0.0.1') self.assertEqual(service['environment']['CODEBUDDY_IMPORT_DIR'], '/data/auth/imports') diff --git a/tests/test_inference_resources.py b/tests/test_inference_resources.py new file mode 100644 index 0000000..a559969 --- /dev/null +++ b/tests/test_inference_resources.py @@ -0,0 +1,417 @@ +"""Verify bounded connection reuse, identity isolation and cancellation-safe account leases.""" +import sys +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +import asyncio +from concurrent.futures import ThreadPoolExecutor +from contextlib import suppress +import json +import os +import threading +from types import SimpleNamespace +import unittest +from unittest.mock import patch + +import httpx +from fastapi import HTTPException + +import converter as gateway +from app import inference_resources as resources, upstream_io +from app.settings import SCHEMA, apply_persisted_settings, resolve_settings, validate_settings +import test_api_flow as fixtures + +URL = "https://copilot.tencent.com/v2/chat/completions" +SSE = fixtures.fixtures.success_sse() +REAL_CLIENT = httpx.AsyncClient + + +class ResourceOwnershipTests(unittest.TestCase): + def test_atomic_capacity_and_idempotent_release(self): + capacity = resources.AccountCapacity() + with ThreadPoolExecutor(max_workers=12) as executor: + leases = list(executor.map(lambda _: capacity.acquire("account", 3, object(), 1), range(50))) + accepted = [lease for lease in leases if lease is not None] + self.assertEqual(len(accepted), 3) + self.assertEqual(capacity.count("account"), 3) + with ThreadPoolExecutor(max_workers=8) as executor: + list(executor.map(lambda lease: lease.release(), accepted * 5)) + self.assertEqual(capacity.count("account"), 0) + self.assertEqual(capacity._counts, {}) + + def test_scope_releases_late_and_duplicate_leases(self): + capacity = resources.AccountCapacity() + scope = resources.RequestResources() + lease = capacity.acquire("account", 1, object(), 1) + scope.add(lease) + scope.add(lease) + scope.close() + scope.close() + self.assertEqual(capacity.count("account"), 0) + late = capacity.acquire("account", 1, object(), 2) + with self.assertRaises(asyncio.CancelledError): + scope.add(late) + self.assertEqual(capacity._counts, {}) + + def test_unlimited_requests_remain_counted_when_limit_changes(self): + capacity = resources.AccountCapacity() + leases = [capacity.acquire("account", 0, object(), 1) for _ in range(3)] + self.assertIsNone(capacity.acquire("account", 1, object(), 1)) + for lease in leases: + lease.release() + lease = capacity.acquire("account", 1, object(), 2) + self.assertIsInstance(lease, tuple) + self.assertEqual(lease[1], 2) + lease.release() + + +class ClientPoolTests(unittest.IsolatedAsyncioTestCase): + async def collect(self, clients=None, *, url=URL, headers=None): + async with upstream_io.open_backend_stream(url, headers or {}, {}, clients=clients) as response: + return await response.aread() + + async def test_clients_are_reused_without_cookie_or_authorization_bleed(self): + seen, clients = [], [] + def handle(request): + seen.append(request) + return httpx.Response(200, content=SSE, headers={"Set-Cookie": "session=synthetic; Path=/"}) + def create(**kwargs): + limits = kwargs["limits"] + self.assertEqual((limits.max_connections, limits.max_keepalive_connections, limits.keepalive_expiry), (64, 16, 30)) + client = REAL_CLIENT(transport=httpx.MockTransport(handle), **kwargs) + clients.append(client) + return client + pool = resources.UpstreamClients() + try: + with patch.object(httpx, "AsyncClient", side_effect=create): + for identity in ("A", "B"): + await self.collect(pool, headers={"Authorization": "Bearer synthetic-" + identity, + "X-User-Id": identity, "X-Tenant-Id": identity}) + await self.collect(pool, url="https://www.workbuddy.ai/v2/chat/completions") + self.assertEqual(len(clients), 2) + self.assertEqual([request.headers["X-User-Id"] for request in seen[:2]], ["A", "B"]) + self.assertEqual(seen[1].headers["Authorization"], "Bearer synthetic-B") + self.assertNotIn("Authorization", seen[2].headers) + self.assertTrue(all("Cookie" not in request.headers for request in seen)) + self.assertTrue(all(not client.cookies for client in clients)) + self.assertTrue(all(client.trust_env for client in clients)) + self.assertTrue(all(request.extensions["timeout"] == {"connect": 15, "read": 300, "write": 60, "pool": 15} + for request in seen)) + self.assertIsNone(pool.get("https://untrusted.invalid/v2/chat/completions")) + finally: + await pool.aclose() + self.assertTrue(all(client.is_closed for client in clients)) + self.assertIsNone(pool.get(URL)) + await pool.aclose() + + async def test_only_safe_connection_failure_uses_one_fresh_client_retry(self): + for failure in (httpx.ConnectError, httpx.ConnectTimeout, httpx.ReadError, + httpx.ReadTimeout, httpx.WriteError, httpx.WriteTimeout, httpx.PoolTimeout): + with self.subTest(failure=failure.__name__): + seen, made = [], [] + def handle(request): + seen.append(request) + if len(seen) == 1: + raise failure("synthetic transport failure") + return httpx.Response(200, content=SSE) + def create(**kwargs): + client = REAL_CLIENT(transport=httpx.MockTransport(handle), **kwargs) + made.append(client) + return client + pool = resources.UpstreamClients() + try: + with patch.object(httpx, "AsyncClient", side_effect=create): + if failure in (httpx.ConnectError, httpx.ConnectTimeout): + self.assertEqual(await self.collect(pool), SSE) + self.assertEqual(len(made), 2) + self.assertTrue(made[1].is_closed) + else: + with self.assertRaises(failure): + await self.collect(pool) + self.assertEqual(len(made), 1) + self.assertFalse(made[0].is_closed) + finally: + await pool.aclose() + + async def test_cancellation_closes_response_but_preserves_pool_for_next_request(self): + entered, closed = asyncio.Event(), asyncio.Event() + class HangingBody(httpx.AsyncByteStream): + async def __aiter__(self): + entered.set() + await asyncio.Event().wait() + yield b"" + async def aclose(self): + closed.set() + calls = [] + def handle(request): + calls.append(request) + return httpx.Response(200, stream=HangingBody()) if len(calls) == 1 else httpx.Response(200, content=SSE) + pool = resources.UpstreamClients() + try: + with patch.object(httpx, "AsyncClient", side_effect=lambda **kw: REAL_CLIENT(transport=httpx.MockTransport(handle), **kw)): + task = asyncio.create_task(self.collect(pool)) + await asyncio.wait_for(entered.wait(), 1) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + self.assertTrue(closed.is_set()) + self.assertEqual(await self.collect(pool), SSE) + self.assertEqual(len(pool._clients), 1) + finally: + await pool.aclose() + + async def test_real_http_connections_are_reused_only_when_enabled(self): + accepted, writers, handlers = [], [], [] + async def handle(reader, writer): + accepted.append(writer) + writers.append(writer) + handlers.append(asyncio.current_task()) + try: + while True: + headers = await reader.readuntil(b"\r\n\r\n") + length = next(int(line.split(b":", 1)[1]) for line in headers.split(b"\r\n") + if line.lower().startswith(b"content-length:")) + await reader.readexactly(length) + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: " + str(len(SSE)).encode() + b"\r\n\r\n" + SSE) + await writer.drain() + except (asyncio.IncompleteReadError, ConnectionError): + pass + finally: + writer.close() + await writer.wait_closed() + server = await asyncio.start_server(handle, "127.0.0.1", 0) + base = f"http://127.0.0.1:{server.sockets[0].getsockname()[1]}" + try: + with patch.dict(resources.PROFILE_ENDPOINTS, {"fixture": base}, clear=True), patch.dict(os.environ, {"NO_PROXY": "127.0.0.1"}): + pool = resources.UpstreamClients() + try: + for _ in range(2): + self.assertEqual(await self.collect(pool, url=base), SSE) + self.assertEqual(len(accepted), 1) + finally: + await pool.aclose() + for _ in range(2): + self.assertEqual(await self.collect(url=base), SSE) + self.assertEqual(len(accepted), 3) + finally: + server.close() + await server.wait_closed() + for writer in writers: + writer.close() + await asyncio.wait_for(asyncio.gather(*handlers), 2) + + async def test_lifespan_closes_clients_and_never_reuses_a_previous_owner(self): + async with resources.inference_lifespan(None) as state: + first = state["upstream_clients"] + self.assertIsInstance(first, resources.UpstreamClients) + self.assertTrue(first._closed) + async with resources.inference_lifespan(None) as state: + self.assertIsNot(first, state["upstream_clients"]) + + +class PoolCapacityTests(fixtures.GatewayFixture, unittest.TestCase): + def setUp(self): + super().setUp() + self.enterContext(patch.dict(gateway.CONFIG, {"max_inflight_per_account": 1})) + self.fx.configure(profiles=("cn-cli", "cn-work")) + + def acquire(self, session="same-session"): + picked = self.fx.pool.headers_for(session, "shared-model", with_generation=True, with_capacity=True) + self.assertIsNotNone(picked) + lease, _ = picked + self.addCleanup(lease.release) + return lease + + def test_busy_sticky_and_early_expiry_account_spills_within_same_cost_tier(self): + with patch.object(self.fx.pool, "_expiry_rank", side_effect=lambda e: (False, 1 if e["uid"] == "cn-cli" else 2)): + first, second = self.acquire(), self.acquire() + self.assertEqual(first[0], self.fx.entries["cn-cli"]["cm"]) + self.assertEqual(second[0], self.fx.entries["cn-work"]["cm"]) + with self.assertRaises(HTTPException) as error: + self.acquire() + self.assertEqual(error.exception.status_code, 503) + self.assertEqual(error.exception.headers["Retry-After"], "3") + self.assertEqual(error.exception.detail["error"]["code"], "credential_concurrency_limit") + first.release() + self.assertIs(self.acquire()[0], first[0]) + + def test_full_free_tier_does_not_fall_through_to_paid_accounts(self): + self.fx.account_catalogs({ + "cn-cli": [fixtures.fixtures.model("shared-model", "x0.00")], + "cn-work": [fixtures.fixtures.model("shared-model", "x1.00")]}) + self.acquire() + with self.assertRaises(HTTPException) as error: + self.acquire("different-session") + self.assertEqual(error.exception.status_code, 503) + self.assertEqual(self.fx.pool._capacity.count(self.fx.entries["cn-work"]["account_key"]), 0) + + def test_strict_binding_remains_strict_when_its_account_is_full(self): + identity = self.fx.entries["cn-cli"]["account_key"] + store = SimpleNamespace(snapshot=lambda: {"revision": 1, "credentials": {}, + "models": {"shared-model": {"credential_ids": [identity]}}}) + with patch.dict(gateway.CONFIG, {"control_store": store}): + self.acquire() + with self.assertRaises(HTTPException) as error: + self.acquire("different-session") + self.assertEqual(error.exception.detail["error"]["code"], "credential_concurrency_limit") + + def test_atomic_header_reservations_never_overbook_accounts(self): + def attempt(index): + try: + return self.fx.pool.headers_for(str(index), "shared-model", with_generation=True, with_capacity=True)[0] + except HTTPException as error: + self.assertEqual(error.status_code, 503) + return None + with ThreadPoolExecutor(max_workers=12) as executor: + leases = [lease for lease in executor.map(attempt, range(30)) if lease is not None] + try: + self.assertEqual(len(leases), 2) + rows = self.fx.pool.snapshot() + self.assertTrue(all(row["in_flight"] == row["max_in_flight"] == 1 for row in rows)) + finally: + for lease in leases: + lease.release() + self.assertEqual(self.fx.pool._capacity._counts, {}) + + def test_reimport_or_delete_does_not_reset_existing_account_capacity(self): + self.fx.configure(profiles=("cn-cli",)) + lease = self.acquire() + cm = lease[0] + cm.invalidate() + self.fx.pool.reload([cm.path]) + with self.assertRaises(HTTPException): + self.acquire() + content = cm.path.read_bytes() + self.fx.pool.remove_file(cm.path.name) + cm.path.write_bytes(content) + self.fx.pool.reload([cm.path]) + with self.assertRaises(HTTPException): + self.acquire() + lease.release() + self.assertIsNotNone(self.acquire()) + + +class RequestCapacityTests(fixtures.GatewayFixture, unittest.IsolatedAsyncioTestCase): + def setUp(self): + super().setUp() + self.enterContext(patch.dict(gateway.CONFIG, {"max_inflight_per_account": 1})) + self.fx.configure(profiles=("cn-cli",)) + + async def test_gateway_reuses_clients_across_protocols_and_same_origin_accounts(self): + self.fx.add_account("account-B", "cn-cli") + self.fx.allowed_profiles.add("account-B") + self.fx.configure(profiles=("cn-cli", "account-B")) + factory = httpx.AsyncClient + before = factory.call_count + with patch.dict(gateway.CONFIG, {"upstream_keepalive": True}): + for route in fixtures.fixtures.GENERATIONS: + for stream in (False, True): + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(self.fx.pool._capacity._counts, {}) + self.assertEqual(factory.call_count - before, 1) + self.assertEqual({request.headers["X-User-Id"] for request in self.fx.requests}, {"cn-cli", "account-B"}) + + async def test_disabled_keepalive_preserves_fresh_client_behavior(self): + factory = httpx.AsyncClient + before = factory.call_count + for _ in range(2): + response = self.fx.client.post("/v1/responses", json=self.fx.payload("responses")) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(factory.call_count - before, 2) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + + async def test_full_account_rejects_then_recovers_after_cancellation(self): + for route in fixtures.fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(route=route, stream=stream): + entered, closed = asyncio.Event(), asyncio.Event() + class HangingBody(httpx.AsyncByteStream): + async def __aiter__(self): + entered.set() + await asyncio.Event().wait() + yield b"" + async def aclose(self): + closed.set() + with self.responder(lambda request: httpx.Response(200, stream=HangingBody())): + async with REAL_CLIENT(transport=httpx.ASGITransport(app=gateway.app), base_url="http://test") as client: + task = asyncio.create_task(client.post("/v1/" + route, json=self.fx.payload(route, stream=stream))) + try: + await asyncio.wait_for(entered.wait(), 2) + second = await client.post("/v1/" + route, json=self.fx.payload(route, stream=stream)) + self.assertEqual(second.status_code, 503, second.text) + self.assertEqual(second.headers["Retry-After"], "3") + self.assertEqual(second.json()["error"]["code"], "credential_concurrency_limit") + finally: + task.cancel() + with suppress(asyncio.CancelledError): + await task + self.assertTrue(closed.is_set()) + self.assertEqual(self.fx.pool._capacity._counts, {}) + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + async def test_late_worker_reservation_after_disconnect_is_released_without_upstream(self): + original = gateway._route_chat + entered, release, finished = threading.Event(), threading.Event(), threading.Event() + def delayed(*args, **kwargs): + entered.set() + release.wait(3) + try: + return original(*args, **kwargs) + finally: + finished.set() + with patch.object(gateway, "_route_chat", side_effect=delayed): + async with REAL_CLIENT(transport=httpx.ASGITransport(app=gateway.app), base_url="http://test") as client: + task = asyncio.create_task(client.post("/v1/responses", json=self.fx.payload("responses"))) + try: + self.assertTrue(await asyncio.to_thread(entered.wait, 1)) + task.cancel() + with self.assertRaises(asyncio.CancelledError): + await task + finally: + release.set() + self.assertTrue(await asyncio.to_thread(finished.wait, 2)) + self.assertEqual(self.fx.pool._capacity._counts, {}) + self.assertEqual(self.fx.requests, []) + + async def test_failover_releases_previous_capacity_before_next_attempt(self): + for route in fixtures.fixtures.GENERATIONS: + for stream in (False, True): + with self.subTest(route=route, stream=stream): + self.fx.configure(profiles=("cn-cli", "cn-work")) + seen = [] + def respond(request): + seen.append(request) + self.assertEqual(sum(self.fx.pool._capacity._counts.values()), 1) + return httpx.Response(503, json={"error": {"message": "synthetic failure"}}) if len(seen) == 1 else httpx.Response(200, content=SSE) + with self.responder(respond), patch.dict(gateway.CONFIG, {"failover_max": 1}): + response = self.fx.client.post("/v1/" + route, json=self.fx.payload(route, stream=stream)) + self.assertEqual(response.status_code, 200, response.text) + self.assertEqual(len(seen), 2) + self.assertEqual(self.fx.pool._capacity._counts, {}) + + +class ConfigurationTests(unittest.TestCase): + def test_defaults_precedence_validation_and_ui_modes(self): + self.assertFalse(SCHEMA["upstream_keepalive"]["default"]) + self.assertEqual(SCHEMA["max_inflight_per_account"]["default"], 0) + store = SimpleNamespace(snapshot=lambda: {"settings": {"upstream_keepalive": True, "max_inflight_per_account": 2}}) + env = {"CODEBUDDY2API_UPSTREAM_KEEPALIVE": "false", "CODEBUDDY2API_MAX_INFLIGHT_PER_ACCOUNT": "3"} + config = {"control_store": store, "max_inflight_per_account": 1} + apply_persisted_settings(config, explicit=("max_inflight_per_account",), environ=env) + self.assertEqual((config["upstream_keepalive"], config["max_inflight_per_account"]), (False, 1)) + items = {item["key"]: item for item in resolve_settings(config)} + self.assertEqual(items["upstream_keepalive"]["mode"], "restart") + self.assertTrue(items["upstream_keepalive"]["locked"]) + self.assertEqual(items["max_inflight_per_account"]["mode"], "hot") + for values in ({"max_inflight_per_account": -1}, {"max_inflight_per_account": True}, {"upstream_keepalive": "true"}): + with self.assertRaises(ValueError): + validate_settings(values) + + +if __name__ == "__main__": + unittest.main()