splitagent 0.0.3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- splitagent/__init__.py +8 -0
- splitagent/__main__.py +6 -0
- splitagent/agents/__init__.py +10 -0
- splitagent/agents/base.py +477 -0
- splitagent/agents/blue.py +57 -0
- splitagent/agents/chat.py +60 -0
- splitagent/agents/prompts.py +462 -0
- splitagent/agents/red.py +75 -0
- splitagent/cli.py +701 -0
- splitagent/config.py +697 -0
- splitagent/core/__init__.py +19 -0
- splitagent/core/bus.py +62 -0
- splitagent/core/context.py +587 -0
- splitagent/core/context_manager.py +381 -0
- splitagent/core/engine.py +424 -0
- splitagent/core/models.py +310 -0
- splitagent/core/proc.py +73 -0
- splitagent/core/sandbox.py +184 -0
- splitagent/core/toolbox.py +520 -0
- splitagent/core/workspace.py +420 -0
- splitagent/desktop/__init__.py +7 -0
- splitagent/desktop/api.py +525 -0
- splitagent/desktop/app.py +1131 -0
- splitagent/desktop/web/app.js +3067 -0
- splitagent/desktop/web/assets/Inter.ttf +0 -0
- splitagent/desktop/web/assets/JetBrainsMonoNerdFontMono-Regular.woff2 +0 -0
- splitagent/desktop/web/index.html +760 -0
- splitagent/desktop/web/styles.css +1612 -0
- splitagent/errors.py +27 -0
- splitagent/llm/__init__.py +8 -0
- splitagent/llm/client.py +488 -0
- splitagent/llm/types.py +172 -0
- splitagent/report/__init__.py +9 -0
- splitagent/report/cvss.py +93 -0
- splitagent/report/generator.py +733 -0
- splitagent/tools/__init__.py +8 -0
- splitagent/tools/base.py +135 -0
- splitagent/tools/defense.py +475 -0
- splitagent/tools/exploit.py +318 -0
- splitagent/tools/http_pool.py +109 -0
- splitagent/tools/knowledge.py +376 -0
- splitagent/tools/recon.py +182 -0
- splitagent/tools/registry.py +62 -0
- splitagent/tools/validate.py +908 -0
- splitagent/tools/web.py +386 -0
- splitagent/tools/workspace_tools.py +411 -0
- splitagent/ui/__init__.py +5 -0
- splitagent/ui/app.py +389 -0
- splitagent/ui/stream.py +234 -0
- splitagent/ui/theme.py +72 -0
- splitagent-0.0.3.dist-info/METADATA +987 -0
- splitagent-0.0.3.dist-info/RECORD +56 -0
- splitagent-0.0.3.dist-info/WHEEL +5 -0
- splitagent-0.0.3.dist-info/entry_points.txt +2 -0
- splitagent-0.0.3.dist-info/licenses/LICENSE +21 -0
- splitagent-0.0.3.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,318 @@
|
|
|
1
|
+
"""Controlled exploitation probes.
|
|
2
|
+
|
|
3
|
+
Every probe here is non-destructive: it injects benign markers and inspects
|
|
4
|
+
the response for evidence. No tool deletes data, writes files or launches a
|
|
5
|
+
shell on the target. When ``safe_mode`` is enabled the probes stop at the
|
|
6
|
+
detection stage and never chain into exploitation.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from typing import Any
|
|
12
|
+
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
|
|
13
|
+
|
|
14
|
+
import httpx
|
|
15
|
+
|
|
16
|
+
from splitagent.tools.base import Tool, ToolContext
|
|
17
|
+
from splitagent.tools.http_pool import get_client
|
|
18
|
+
|
|
19
|
+
SQL_ERRORS = (
|
|
20
|
+
"you have an error in your sql syntax",
|
|
21
|
+
"warning: mysql",
|
|
22
|
+
"unclosed quotation mark",
|
|
23
|
+
"quoted string not properly terminated",
|
|
24
|
+
"sqlstate",
|
|
25
|
+
"pg_query",
|
|
26
|
+
"psql:",
|
|
27
|
+
"syntax error at or near",
|
|
28
|
+
"ora-01756",
|
|
29
|
+
"ora-00933",
|
|
30
|
+
"microsoft odbc",
|
|
31
|
+
"odbc sql server driver",
|
|
32
|
+
"native client",
|
|
33
|
+
"sqlite3::",
|
|
34
|
+
"jdbc",
|
|
35
|
+
"doctrine\\query",
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
CANARY = "splitagentcanary"
|
|
39
|
+
XSS_MARKER = f"<svg/onload=alert('{CANARY}')>"
|
|
40
|
+
CMD_MARKER = f"splitagent_{CANARY}"
|
|
41
|
+
TRAVERSAL_PAYLOADS = (
|
|
42
|
+
"../../../../../../etc/passwd",
|
|
43
|
+
"..%2f..%2f..%2f..%2fetc%2fpasswd",
|
|
44
|
+
"....//....//....//etc/passwd",
|
|
45
|
+
"/etc/passwd",
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _inject(url: str, parameter: str, value: str) -> str:
|
|
50
|
+
parsed = urlparse(url)
|
|
51
|
+
query = dict(parse_qsl(parsed.query, keep_blank_values=True))
|
|
52
|
+
query[parameter] = value
|
|
53
|
+
return urlunparse(parsed._replace(query=urlencode(query)))
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
async def _send(
|
|
57
|
+
ctx: ToolContext, url: str, method: str = "GET", data: dict[str, Any] | None = None
|
|
58
|
+
) -> httpx.Response:
|
|
59
|
+
ctx.check_scope(url)
|
|
60
|
+
client = await get_client(follow=False, verify=False, timeout=15.0, headers=ctx.auth_headers())
|
|
61
|
+
return await client.request(method.upper(), url, data=data)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def _evidence(response: httpx.Response, needles: list[str]) -> str:
|
|
65
|
+
body = response.text[:20000]
|
|
66
|
+
for needle in needles:
|
|
67
|
+
index = body.lower().find(needle.lower())
|
|
68
|
+
if index >= 0:
|
|
69
|
+
start = max(0, index - 120)
|
|
70
|
+
return body[start : index + len(needle) + 120].replace("\n", " ")
|
|
71
|
+
return body[:400].replace("\n", " ")
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
async def test_sqli(
|
|
75
|
+
ctx: ToolContext, url: str, parameter: str, method: str = "GET"
|
|
76
|
+
) -> dict[str, Any]:
|
|
77
|
+
baseline = await _send(ctx, url, method)
|
|
78
|
+
findings: list[str] = []
|
|
79
|
+
for payload in ("'", '"', "')", "' OR '1'='1"):
|
|
80
|
+
probe_url = _inject(url, parameter, payload)
|
|
81
|
+
response = await _send(ctx, probe_url, method)
|
|
82
|
+
body = response.text.lower()
|
|
83
|
+
for signature in SQL_ERRORS:
|
|
84
|
+
if signature in body:
|
|
85
|
+
findings.append(f"payload={payload!r} signature={signature!r}")
|
|
86
|
+
break
|
|
87
|
+
vulnerable = bool(findings)
|
|
88
|
+
return {
|
|
89
|
+
"vulnerable": vulnerable,
|
|
90
|
+
"test": "sql_injection",
|
|
91
|
+
"parameter": parameter,
|
|
92
|
+
"baseline_status": baseline.status_code,
|
|
93
|
+
"evidence": findings or "no SQL error signatures observed",
|
|
94
|
+
"confidence": "high" if vulnerable else "low",
|
|
95
|
+
}
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
async def test_xss(
|
|
99
|
+
ctx: ToolContext, url: str, parameter: str, method: str = "GET"
|
|
100
|
+
) -> dict[str, Any]:
|
|
101
|
+
probe_url = _inject(url, parameter, XSS_MARKER)
|
|
102
|
+
response = await _send(ctx, probe_url, method)
|
|
103
|
+
reflected = XSS_MARKER.lower() in response.text.lower()
|
|
104
|
+
content_type = response.headers.get("content-type", "")
|
|
105
|
+
html_context = "html" in content_type.lower() or "<html" in response.text[:500].lower()
|
|
106
|
+
return {
|
|
107
|
+
"vulnerable": bool(reflected and html_context),
|
|
108
|
+
"test": "reflected_xss",
|
|
109
|
+
"parameter": parameter,
|
|
110
|
+
"reflected": reflected,
|
|
111
|
+
"content_type": content_type,
|
|
112
|
+
"evidence": _evidence(response, [CANARY]) if reflected else "marker not reflected",
|
|
113
|
+
"confidence": "high" if (reflected and html_context) else "low",
|
|
114
|
+
}
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
async def test_path_traversal(
|
|
118
|
+
ctx: ToolContext, url: str, parameter: str, method: str = "GET"
|
|
119
|
+
) -> dict[str, Any]:
|
|
120
|
+
for payload in TRAVERSAL_PAYLOADS:
|
|
121
|
+
probe_url = _inject(url, parameter, payload)
|
|
122
|
+
response = await _send(ctx, probe_url, method)
|
|
123
|
+
if "root:x:0:0" in response.text or "daemon:x:" in response.text:
|
|
124
|
+
return {
|
|
125
|
+
"vulnerable": True,
|
|
126
|
+
"test": "path_traversal",
|
|
127
|
+
"parameter": parameter,
|
|
128
|
+
"payload": payload,
|
|
129
|
+
"evidence": _evidence(response, ["root:x:0:0", "daemon:x:"]),
|
|
130
|
+
"confidence": "high",
|
|
131
|
+
}
|
|
132
|
+
return {
|
|
133
|
+
"vulnerable": False,
|
|
134
|
+
"test": "path_traversal",
|
|
135
|
+
"parameter": parameter,
|
|
136
|
+
"evidence": "no file-disclosure signature observed",
|
|
137
|
+
"confidence": "low",
|
|
138
|
+
}
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
async def test_open_redirect(
|
|
142
|
+
ctx: ToolContext, url: str, parameter: str, method: str = "GET"
|
|
143
|
+
) -> dict[str, Any]:
|
|
144
|
+
canary_url = "https://example.com/splitagent-redirect-canary"
|
|
145
|
+
probe_url = _inject(url, parameter, canary_url)
|
|
146
|
+
response = await _send(ctx, probe_url, method)
|
|
147
|
+
location = response.headers.get("location", "")
|
|
148
|
+
redirected = "example.com/splitagent-redirect-canary" in location
|
|
149
|
+
return {
|
|
150
|
+
"vulnerable": redirected,
|
|
151
|
+
"test": "open_redirect",
|
|
152
|
+
"parameter": parameter,
|
|
153
|
+
"status": response.status_code,
|
|
154
|
+
"location": location,
|
|
155
|
+
"evidence": f"Location: {location}" if redirected else "no external redirect",
|
|
156
|
+
"confidence": "high" if redirected else "low",
|
|
157
|
+
}
|
|
158
|
+
|
|
159
|
+
|
|
160
|
+
async def test_command_injection(
|
|
161
|
+
ctx: ToolContext, url: str, parameter: str, method: str = "GET"
|
|
162
|
+
) -> dict[str, Any]:
|
|
163
|
+
payload = f";echo {CMD_MARKER}"
|
|
164
|
+
probe_url = _inject(url, parameter, payload)
|
|
165
|
+
response = await _send(ctx, probe_url, method)
|
|
166
|
+
injected = CMD_MARKER in response.text
|
|
167
|
+
return {
|
|
168
|
+
"vulnerable": injected,
|
|
169
|
+
"test": "command_injection",
|
|
170
|
+
"parameter": parameter,
|
|
171
|
+
"evidence": _evidence(response, [CMD_MARKER]) if injected else "marker not found",
|
|
172
|
+
"confidence": "high" if injected else "low",
|
|
173
|
+
}
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
async def test_cors(ctx: ToolContext, url: str) -> dict[str, Any]:
|
|
177
|
+
origin = "https://evil.splitagent.test"
|
|
178
|
+
ctx.check_scope(url)
|
|
179
|
+
client = await get_client(follow=False, verify=False, timeout=10.0, headers=ctx.auth_headers())
|
|
180
|
+
response = await client.get(url, headers={"Origin": origin})
|
|
181
|
+
allow_origin = response.headers.get("access-control-allow-origin", "")
|
|
182
|
+
allow_creds = response.headers.get("access-control-allow-credentials", "")
|
|
183
|
+
vulnerable = allow_origin in ("*", origin)
|
|
184
|
+
return {
|
|
185
|
+
"vulnerable": vulnerable,
|
|
186
|
+
"test": "cors_misconfiguration",
|
|
187
|
+
"allow_origin": allow_origin,
|
|
188
|
+
"allow_credentials": allow_creds,
|
|
189
|
+
"evidence": f"ACAO={allow_origin!r} ACAC={allow_creds!r}",
|
|
190
|
+
"confidence": "high" if vulnerable else "low",
|
|
191
|
+
}
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
async def test_http_methods(ctx: ToolContext, url: str) -> dict[str, Any]:
|
|
195
|
+
ctx.check_scope(url)
|
|
196
|
+
allowed: list[str] = []
|
|
197
|
+
client = await get_client(follow=False, verify=False, timeout=10.0, headers=ctx.auth_headers())
|
|
198
|
+
for method in ("GET", "POST", "PUT", "DELETE", "TRACE", "OPTIONS", "PATCH"):
|
|
199
|
+
try:
|
|
200
|
+
response = await client.request(method, url)
|
|
201
|
+
except httpx.HTTPError:
|
|
202
|
+
continue
|
|
203
|
+
if response.status_code < 400 or response.status_code == 405:
|
|
204
|
+
allowed.append(f"{method}:{response.status_code}")
|
|
205
|
+
dangerous = [m for m in allowed if m.split(":")[0] in ("PUT", "DELETE", "TRACE")]
|
|
206
|
+
return {
|
|
207
|
+
"vulnerable": bool(dangerous),
|
|
208
|
+
"test": "http_methods",
|
|
209
|
+
"allowed": allowed,
|
|
210
|
+
"dangerous": dangerous,
|
|
211
|
+
"evidence": ", ".join(allowed),
|
|
212
|
+
"confidence": "medium" if dangerous else "low",
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
async def test_directory_listing(ctx: ToolContext, url: str) -> dict[str, Any]:
|
|
217
|
+
response = await _send(ctx, url)
|
|
218
|
+
body = response.text.lower()
|
|
219
|
+
listing = "index of /" in body and "<a href=" in body
|
|
220
|
+
return {
|
|
221
|
+
"vulnerable": listing,
|
|
222
|
+
"test": "directory_listing",
|
|
223
|
+
"status": response.status_code,
|
|
224
|
+
"evidence": _evidence(response, ["index of /"]) if listing else "no listing",
|
|
225
|
+
"confidence": "high" if listing else "low",
|
|
226
|
+
}
|
|
227
|
+
|
|
228
|
+
|
|
229
|
+
def red_exploit_tools(ctx: ToolContext) -> list[Tool]:
|
|
230
|
+
def _params(extra: dict[str, Any]) -> dict[str, Any]:
|
|
231
|
+
base = {
|
|
232
|
+
"url": {"type": "string"},
|
|
233
|
+
"parameter": {"type": "string"},
|
|
234
|
+
"method": {"type": "string", "default": "GET"},
|
|
235
|
+
}
|
|
236
|
+
base.update(extra)
|
|
237
|
+
return {"type": "object", "properties": base, "required": ["url", "parameter"]}
|
|
238
|
+
|
|
239
|
+
return [
|
|
240
|
+
Tool(
|
|
241
|
+
name="test_sql_injection",
|
|
242
|
+
description=(
|
|
243
|
+
"Non-destructive SQL injection detection: injects quote payloads "
|
|
244
|
+
"and looks for database error signatures."
|
|
245
|
+
),
|
|
246
|
+
parameters=_params({}),
|
|
247
|
+
func=lambda url, parameter, method="GET": test_sqli(ctx, url, parameter, method),
|
|
248
|
+
scope="red",
|
|
249
|
+
),
|
|
250
|
+
Tool(
|
|
251
|
+
name="test_xss",
|
|
252
|
+
description=("Reflected XSS detection using a unique marker in the parameter."),
|
|
253
|
+
parameters=_params({}),
|
|
254
|
+
func=lambda url, parameter, method="GET": test_xss(ctx, url, parameter, method),
|
|
255
|
+
scope="red",
|
|
256
|
+
),
|
|
257
|
+
Tool(
|
|
258
|
+
name="test_path_traversal",
|
|
259
|
+
description="Detect file disclosure via path traversal payloads.",
|
|
260
|
+
parameters=_params({}),
|
|
261
|
+
func=lambda url, parameter, method="GET": test_path_traversal(
|
|
262
|
+
ctx, url, parameter, method
|
|
263
|
+
),
|
|
264
|
+
scope="red",
|
|
265
|
+
),
|
|
266
|
+
Tool(
|
|
267
|
+
name="test_open_redirect",
|
|
268
|
+
description="Detect open redirect by injecting an external canary URL.",
|
|
269
|
+
parameters=_params({}),
|
|
270
|
+
func=lambda url, parameter, method="GET": test_open_redirect(
|
|
271
|
+
ctx, url, parameter, method
|
|
272
|
+
),
|
|
273
|
+
scope="red",
|
|
274
|
+
),
|
|
275
|
+
Tool(
|
|
276
|
+
name="test_command_injection",
|
|
277
|
+
description=("Detect OS command injection using a benign echo marker."),
|
|
278
|
+
parameters=_params({}),
|
|
279
|
+
func=lambda url, parameter, method="GET": test_command_injection(
|
|
280
|
+
ctx, url, parameter, method
|
|
281
|
+
),
|
|
282
|
+
scope="red",
|
|
283
|
+
dangerous=True,
|
|
284
|
+
),
|
|
285
|
+
Tool(
|
|
286
|
+
name="test_cors",
|
|
287
|
+
description="Check for permissive CORS configuration.",
|
|
288
|
+
parameters={
|
|
289
|
+
"type": "object",
|
|
290
|
+
"properties": {"url": {"type": "string"}},
|
|
291
|
+
"required": ["url"],
|
|
292
|
+
},
|
|
293
|
+
func=lambda url: test_cors(ctx, url),
|
|
294
|
+
scope="red",
|
|
295
|
+
),
|
|
296
|
+
Tool(
|
|
297
|
+
name="test_http_methods",
|
|
298
|
+
description="Enumerate supported HTTP methods and flag dangerous ones.",
|
|
299
|
+
parameters={
|
|
300
|
+
"type": "object",
|
|
301
|
+
"properties": {"url": {"type": "string"}},
|
|
302
|
+
"required": ["url"],
|
|
303
|
+
},
|
|
304
|
+
func=lambda url: test_http_methods(ctx, url),
|
|
305
|
+
scope="red",
|
|
306
|
+
),
|
|
307
|
+
Tool(
|
|
308
|
+
name="test_directory_listing",
|
|
309
|
+
description="Check whether a directory exposes an automatic index.",
|
|
310
|
+
parameters={
|
|
311
|
+
"type": "object",
|
|
312
|
+
"properties": {"url": {"type": "string"}},
|
|
313
|
+
"required": ["url"],
|
|
314
|
+
},
|
|
315
|
+
func=lambda url: test_directory_listing(ctx, url),
|
|
316
|
+
scope="red",
|
|
317
|
+
),
|
|
318
|
+
]
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
"""A shared HTTP connection pool for the reconnaissance tools.
|
|
2
|
+
|
|
3
|
+
Every web tool used to build its own ``httpx.AsyncClient`` per call, which
|
|
4
|
+
means a fresh TCP handshake, a fresh TLS handshake and a fresh DNS lookup for
|
|
5
|
+
every single request. During a crawl or a path sweep that dominates the wall
|
|
6
|
+
clock. One client with keep-alive removes it entirely.
|
|
7
|
+
|
|
8
|
+
The pool is keyed by the flags that affect connection reuse (redirect and
|
|
9
|
+
verification policy), created lazily and closed by the engine at the end of a
|
|
10
|
+
run.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import asyncio
|
|
16
|
+
from typing import Any
|
|
17
|
+
|
|
18
|
+
import httpx
|
|
19
|
+
|
|
20
|
+
_CLIENTS: dict[tuple[Any, ...], httpx.AsyncClient] = {}
|
|
21
|
+
_LOCK = asyncio.Lock()
|
|
22
|
+
|
|
23
|
+
DEFAULT_LIMITS = httpx.Limits(
|
|
24
|
+
max_connections=64,
|
|
25
|
+
max_keepalive_connections=32,
|
|
26
|
+
keepalive_expiry=30.0,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _http2_available() -> bool:
|
|
31
|
+
"""HTTP/2 needs the optional ``h2`` package; degrade silently without it."""
|
|
32
|
+
try:
|
|
33
|
+
import h2 # noqa: F401
|
|
34
|
+
except ImportError:
|
|
35
|
+
return False
|
|
36
|
+
return True
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
HTTP2 = _http2_available()
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _key(
|
|
43
|
+
follow: bool, verify: bool, timeout: float, headers: dict[str, str] | None
|
|
44
|
+
) -> tuple[Any, ...]:
|
|
45
|
+
# Headers participate in the key so authenticated and anonymous traffic
|
|
46
|
+
# never share a connection pool (and therefore never leak a cookie).
|
|
47
|
+
header_key = tuple(sorted((headers or {}).items()))
|
|
48
|
+
return (follow, verify, timeout, header_key)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
async def get_client(
|
|
52
|
+
*,
|
|
53
|
+
follow: bool = True,
|
|
54
|
+
verify: bool = False,
|
|
55
|
+
timeout: float = 15.0,
|
|
56
|
+
headers: dict[str, str] | None = None,
|
|
57
|
+
) -> httpx.AsyncClient:
|
|
58
|
+
"""Return a pooled client for the given transport policy.
|
|
59
|
+
|
|
60
|
+
The headers become the client's defaults. Callers must NOT pass the same
|
|
61
|
+
headers per-request: httpx appends request headers to the client defaults,
|
|
62
|
+
which would send each one twice and the server would reject the request as
|
|
63
|
+
having conflicting headers. Use :func:`request_headers` to get the extras
|
|
64
|
+
that still need to be added (a per-call override such as ``Origin``).
|
|
65
|
+
"""
|
|
66
|
+
key = _key(follow, verify, timeout, headers)
|
|
67
|
+
client = _CLIENTS.get(key)
|
|
68
|
+
if client is not None and not client.is_closed:
|
|
69
|
+
return client
|
|
70
|
+
async with _LOCK:
|
|
71
|
+
client = _CLIENTS.get(key)
|
|
72
|
+
if client is not None and not client.is_closed:
|
|
73
|
+
return client
|
|
74
|
+
client = httpx.AsyncClient(
|
|
75
|
+
timeout=timeout,
|
|
76
|
+
follow_redirects=follow,
|
|
77
|
+
verify=verify,
|
|
78
|
+
headers=dict(headers or {}),
|
|
79
|
+
limits=DEFAULT_LIMITS,
|
|
80
|
+
http2=HTTP2,
|
|
81
|
+
)
|
|
82
|
+
_CLIENTS[key] = client
|
|
83
|
+
return client
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def request_headers(
|
|
87
|
+
client: httpx.AsyncClient, overrides: dict[str, str] | None = None
|
|
88
|
+
) -> dict[str, str]:
|
|
89
|
+
"""Return only the headers a caller must add on top of the client defaults.
|
|
90
|
+
|
|
91
|
+
Anything already present as a client default is dropped, so a request never
|
|
92
|
+
carries the same header twice.
|
|
93
|
+
"""
|
|
94
|
+
if not overrides:
|
|
95
|
+
return {}
|
|
96
|
+
defaults = {name.lower() for name in client.headers}
|
|
97
|
+
return {name: value for name, value in overrides.items() if name.lower() not in defaults}
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
async def aclose_all() -> None:
|
|
101
|
+
"""Close every pooled client (called when a run finishes)."""
|
|
102
|
+
async with _LOCK:
|
|
103
|
+
clients = list(_CLIENTS.values())
|
|
104
|
+
_CLIENTS.clear()
|
|
105
|
+
for client in clients:
|
|
106
|
+
try:
|
|
107
|
+
await client.aclose()
|
|
108
|
+
except Exception:
|
|
109
|
+
pass
|