Repository navigation
Expand file tree
/
Copy pathtest_batch_middleware.py
More file actions
109 lines (84 loc) · 3.42 KB
/
Copy pathtest_batch_middleware.py
File metadata and controls
109 lines (84 loc) · 3.42 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
"""Tests for BatchGatewayMiddleware.__call__ routing — specifically the no-param
content-fetch interception (a stock OpenAI SDK ``files.content()`` works without
the ``custom_llm_provider`` query param)."""
from __future__ import annotations
import pytest
from airlock.batch.middleware import BatchGatewayMiddleware
def _scope(method, path, query=b""):
return {
"type": "http",
"method": method,
"path": path,
"query_string": query,
"headers": [],
}
class _Inner:
"""Fake LiteLLM app: records whether it was reached."""
def __init__(self):
self.called = False
async def __call__(self, scope, receive, send):
self.called = True
await send({"type": "http.response.start", "status": 299, "headers": []})
await send({"type": "http.response.body", "body": b"INNER"})
class _Cap:
def __init__(self):
self.status = None
self.body = b""
async def __call__(self, m):
if m["type"] == "http.response.start":
self.status = m["status"]
elif m["type"] == "http.response.body":
self.body += m.get("body", b"")
async def _receive():
return {"type": "http.request", "body": b"", "more_body": False}
@pytest.fixture(autouse=True)
def _wire(tmp_path, monkeypatch):
monkeypatch.setenv("AIRLOCK_STATE_DIR", str(tmp_path))
monkeypatch.delenv("AIRLOCK_MASTER_KEY", raising=False)
def _stage(file_id, rows):
from airlock.batch import runtime
runtime.write_output(file_id, rows)
class TestNoParamContentInterception:
async def test_gateway_file_is_served_without_param(self):
fid = "file-" + "a" * 32
_stage(fid, [{"custom_id": "r1", "response": {"body": {"ok": 1}}}])
inner = _Inner()
cap = _Cap()
await BatchGatewayMiddleware(inner)(
_scope("GET", f"/v1/files/{fid}/content"), _receive, cap
)
assert inner.called is False # gateway intercepted, no param needed
assert cap.status == 200
assert b"custom_id" in cap.body
async def test_unknown_file_falls_through_to_litellm(self):
inner = _Inner()
cap = _Cap()
await BatchGatewayMiddleware(inner)(
_scope("GET", "/v1/files/file-" + "b" * 32 + "/content"), _receive, cap
)
assert inner.called is True # not a gateway file -> native handler
assert cap.body == b"INNER"
async def test_traversal_id_falls_through_without_fs_use(self):
inner = _Inner()
cap = _Cap()
await BatchGatewayMiddleware(inner)(
_scope("GET", "/v1/files/../../etc/passwd/content"), _receive, cap
)
assert inner.called is True # id fails the strict pattern -> no intercept
async def test_upload_post_still_requires_param(self):
inner = _Inner()
cap = _Cap()
await BatchGatewayMiddleware(inner)(_scope("POST", "/v1/files"), _receive, cap)
assert inner.called is True # uploads are NOT intercepted without the param
async def test_param_path_still_dispatches(self):
fid = "file-" + "c" * 32
_stage(fid, [{"custom_id": "r1", "response": {"body": {}}}])
inner = _Inner()
cap = _Cap()
await BatchGatewayMiddleware(inner)(
_scope("GET", f"/v1/files/{fid}/content", b"custom_llm_provider=vllm"),
_receive,
cap,
)
assert inner.called is False
assert cap.status == 200