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()