daili / test_xai_auth_proxy.py
pjpjq's picture
fix(proxy): 强制 xAI 认证使用多出口代理
bc3a943
Raw
History Blame Contribute Delete
19.6 kB
from __future__ import annotations
import hashlib
import io
import json
import os
import stat
import tempfile
import unittest
from contextlib import redirect_stderr
from pathlib import Path
from typing import Any
from urllib.parse import parse_qs, urlparse
import xai_auth_proxy
PROXIES = (
"socks5h://user:pass@proxy-a.example:1080",
"socks5h://user:pass@proxy-b.example:1080",
"http://proxy-c.example:8080",
)
class FakeResponse:
def __init__(self, body: bytes = b'{"status":"ok"}', status: int = 200) -> None:
self.body = body
self.status = status
self.closed = False
def read(self) -> bytes:
return self.body
def close(self) -> None:
self.closed = True
class FakeManagementClient:
def __init__(self, entries: list[dict[str, Any]], payloads: dict[str, dict[str, Any]]) -> None:
self.entries = entries
self.payloads = payloads
self.downloads: list[str] = []
self.patches: list[tuple[str, str]] = []
self.download_error: Exception | None = None
def list_auth_files(self) -> list[dict[str, Any]]:
return [dict(entry) for entry in self.entries]
def download_auth_file(self, name: str) -> dict[str, Any]:
self.downloads.append(name)
if self.download_error is not None:
error = self.download_error
self.download_error = None
raise error
return dict(self.payloads[name])
def patch_auth_proxy(self, name: str, proxy_url: str) -> None:
self.patches.append((name, proxy_url))
self.payloads[name]["proxy_url"] = proxy_url
class ProxyParsingAndAssignmentTests(unittest.TestCase):
def test_parse_comma_and_newline_separated_urls(self) -> None:
raw = " socks5h://a:1, socks5h://b:2\r\n\n socks5h://a:1 ,http://c:3 "
self.assertEqual(
xai_auth_proxy.parse_proxy_urls(raw),
("socks5h://a:1", "socks5h://b:2", "http://c:3"),
)
def test_sha256_assignment_is_stable_and_uses_all_proxies(self) -> None:
filename = "xai-user@example.com.json"
expected_index = int.from_bytes(
hashlib.sha256(filename.encode("utf-8")).digest(), "big"
) % len(PROXIES)
expected = PROXIES[expected_index]
self.assertEqual(xai_auth_proxy.stable_proxy_for_auth(filename, PROXIES), expected)
self.assertEqual(xai_auth_proxy.stable_proxy_for_file(f"/tmp/{filename}", PROXIES), expected)
assignments = {
xai_auth_proxy.stable_proxy_for_auth(f"xai-account-{index}.json", PROXIES)
for index in range(200)
}
self.assertEqual(assignments, set(PROXIES))
def test_direct_and_malformed_proxy_values_are_rejected(self) -> None:
for value in ("direct", "none", "not-a-proxy", "file:///tmp/socket"):
with self.subTest(value=value):
with self.assertRaisesRegex(ValueError, "invalid xAI proxy URL"):
xai_auth_proxy.parse_proxy_urls(value)
class LocalEnforcementTests(unittest.TestCase):
def test_only_xai_is_updated_atomically_with_mode_0600(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
xai_path = root / "xai-alpha.json"
codex_path = root / "codex-alpha.json"
xai_path.write_text(json.dumps({"type": "XAI", "access_token": "secret"}), encoding="utf-8")
codex_original = b'{ "type": "codex", "proxy_url": "direct", "keep": true }\n'
codex_path.write_bytes(codex_original)
xai_path.chmod(0o644)
codex_path.chmod(0o644)
result = xai_auth_proxy.enforce_local_auths(root, PROXIES)
self.assertEqual(result.scanned, 2)
self.assertEqual(result.xai, 1)
self.assertEqual(result.updated, 1)
xai_payload = json.loads(xai_path.read_text(encoding="utf-8"))
self.assertEqual(
xai_payload["proxy_url"],
xai_auth_proxy.stable_proxy_for_auth(xai_path.name, PROXIES),
)
self.assertEqual(stat.S_IMODE(xai_path.stat().st_mode), 0o600)
self.assertEqual(codex_path.read_bytes(), codex_original)
self.assertEqual(stat.S_IMODE(codex_path.stat().st_mode), 0o644)
self.assertEqual(list(root.glob(".*.tmp")), [])
def test_bad_json_is_never_overwritten(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "xai-broken.json"
original = b'{"type":"xai","access_token":'
path.write_bytes(original)
path.chmod(0o644)
result = xai_auth_proxy.enforce_local_auths(directory, PROXIES)
self.assertEqual(result.invalid_json, 1)
self.assertEqual(result.updated, 0)
self.assertEqual(path.read_bytes(), original)
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o644)
def test_correct_xai_proxy_keeps_content_and_repairs_permission(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "xai-correct.json"
proxy = xai_auth_proxy.stable_proxy_for_auth(path.name, PROXIES)
original = (json.dumps({"type": "xai", "proxy_url": proxy}) + "\n").encode("utf-8")
path.write_bytes(original)
path.chmod(0o644)
result = xai_auth_proxy.enforce_local_auths(directory, PROXIES)
self.assertEqual(result.updated, 0)
self.assertEqual(result.permissions_fixed, 1)
self.assertEqual(path.read_bytes(), original)
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o600)
def test_unwritable_xai_is_quarantined_instead_of_left_direct(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "xai-write-failure.json"
original = b'{"type":"xai","proxy_url":"direct"}'
path.write_bytes(original)
def failing_writer(*args: Any, **kwargs: Any) -> None:
raise OSError("simulated write failure")
result = xai_auth_proxy.enforce_local_auths(
directory,
PROXIES,
writer=failing_writer,
)
self.assertEqual(result.errors, [])
self.assertEqual(result.quarantined, 1)
self.assertFalse(path.exists())
quarantined = list(Path(directory).glob("xai-write-failure.json.xai-proxy-disabled-*"))
self.assertEqual(len(quarantined), 1)
self.assertEqual(quarantined[0].read_bytes(), original)
def test_atomic_writer_fsyncs_and_replaces(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "xai.json"
path.write_text("{}", encoding="utf-8")
fsynced_modes: list[int] = []
replacements: list[tuple[str, str]] = []
def fsync(descriptor: int) -> None:
fsynced_modes.append(os.fstat(descriptor).st_mode)
os.fsync(descriptor)
def replace(source: str, destination: str) -> None:
replacements.append((source, destination))
os.replace(source, destination)
xai_auth_proxy.atomic_write_json(
path,
{"type": "xai", "proxy_url": PROXIES[0]},
replace=replace,
fsync=fsync,
)
self.assertTrue(any(stat.S_ISREG(mode) for mode in fsynced_modes))
self.assertTrue(any(stat.S_ISDIR(mode) for mode in fsynced_modes))
self.assertEqual(len(replacements), 1)
self.assertEqual(replacements[0][1], str(path))
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o600)
self.assertEqual(json.loads(path.read_text(encoding="utf-8"))["type"], "xai")
class GatewayConfigIsolationTests(unittest.TestCase):
def test_global_proxy_is_replaced_without_touching_nested_keys_or_comment(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "config.yaml"
path.write_text(
'host: ""\nproxy-url: "socks5h://127.0.0.1:1080" # global\n'
'provider:\n proxy-url: "provider-specific"\n',
encoding="utf-8",
)
path.chmod(0o640)
changed = xai_auth_proxy.enforce_config_proxy_url(path)
self.assertTrue(changed)
self.assertEqual(
path.read_text(encoding="utf-8"),
'host: ""\nproxy-url: "direct" # global\n'
'provider:\n proxy-url: "provider-specific"\n',
)
self.assertEqual(stat.S_IMODE(path.stat().st_mode), 0o640)
self.assertFalse(xai_auth_proxy.enforce_config_proxy_url(path))
def test_global_proxy_is_appended_when_missing(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "config.yaml"
path.write_text("host: localhost", encoding="utf-8")
self.assertTrue(xai_auth_proxy.enforce_config_proxy_url(path))
self.assertEqual(path.read_text(encoding="utf-8"), 'host: localhost\nproxy-url: "direct"\n')
class ManagementClientTests(unittest.TestCase):
def test_patch_uses_fields_endpoint_and_expected_body(self) -> None:
calls: list[Any] = []
response = FakeResponse()
def opener(request: Any, timeout: float) -> FakeResponse:
calls.append((request, timeout))
return response
client = xai_auth_proxy.ManagementClient(
"http://127.0.0.1:8317/v0/management/",
"management-secret",
timeout=3.5,
opener=opener,
)
client.patch_auth_proxy("xai name@example.com.json", PROXIES[1])
self.assertEqual(len(calls), 1)
request, timeout = calls[0]
self.assertEqual(timeout, 3.5)
self.assertEqual(request.get_method(), "PATCH")
self.assertEqual(request.full_url, "http://127.0.0.1:8317/v0/management/auth-files/fields")
self.assertEqual(
json.loads(request.data),
{"name": "xai name@example.com.json", "proxy_url": PROXIES[1]},
)
self.assertEqual(request.get_header("Authorization"), "Bearer management-secret")
self.assertEqual(request.get_header("X-management-key"), "management-secret")
self.assertTrue(response.closed)
def test_list_and_download_validate_management_payloads(self) -> None:
requests: list[Any] = []
responses = [
FakeResponse(b'{"files":[{"name":"xai-a.json","type":"xai"}]}'),
FakeResponse(b'{"type":"xai","access_token":"secret"}'),
]
def opener(request: Any, timeout: float) -> FakeResponse:
requests.append(request)
return responses.pop(0)
client = xai_auth_proxy.ManagementClient("http://local/v0/management", "key", opener=opener)
self.assertEqual(client.list_auth_files()[0]["name"], "xai-a.json")
self.assertEqual(client.download_auth_file("xai a.json")["type"], "xai")
parsed = urlparse(requests[1].full_url)
self.assertEqual(parsed.path, "/v0/management/auth-files/download")
self.assertEqual(parse_qs(parsed.query), {"name": ["xai a.json"]})
class ManagementEnforcementTests(unittest.TestCase):
def test_only_xai_is_downloaded_and_patched_then_cached(self) -> None:
xai_name = "xai-alpha.json"
codex_name = "codex-alpha.json"
entries = [
{"name": xai_name, "type": "xai", "modtime": "one", "size": 10},
{"name": codex_name, "type": "codex", "modtime": "one", "size": 10},
]
client = FakeManagementClient(
entries,
{
xai_name: {"type": "xai", "proxy_url": "direct"},
codex_name: {"type": "codex", "proxy_url": "direct"},
},
)
cache: dict[str, xai_auth_proxy.AuthFileSignature] = {}
first = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
second = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
expected = xai_auth_proxy.stable_proxy_for_auth(xai_name, PROXIES)
self.assertEqual(first.downloaded, 2)
self.assertEqual(first.patched, 1)
self.assertEqual(client.downloads, [xai_name, xai_name, xai_name])
self.assertEqual(client.patches, [(xai_name, expected)])
self.assertEqual(second.cached, 1)
self.assertEqual(second.downloaded, 1)
self.assertNotIn(codex_name, cache)
def test_changed_modtime_or_size_repairs_overwritten_xai(self) -> None:
name = "xai-overwritten.json"
client = FakeManagementClient(
[{"name": name, "type": "xai", "modtime": "one", "size": 10}],
{name: {"type": "xai", "proxy_url": "direct"}},
)
cache: dict[str, xai_auth_proxy.AuthFileSignature] = {}
xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
client.payloads[name]["proxy_url"] = "http://overwriter.invalid:1"
client.entries[0]["modtime"] = "two"
modtime_change = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
client.payloads[name]["proxy_url"] = "http://overwriter.invalid:2"
client.entries[0]["size"] = 11
size_change = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
self.assertEqual(modtime_change.patched, 1)
self.assertEqual(modtime_change.downloaded, 2)
self.assertEqual(size_change.downloaded, 2)
self.assertEqual(size_change.patched, 1)
self.assertEqual(len(client.patches), 3)
def test_same_signature_overwrite_is_still_repaired(self) -> None:
name = "xai-same-signature.json"
client = FakeManagementClient(
[{"name": name, "type": "xai", "modtime": "same", "size": 10}],
{name: {"type": "xai", "proxy_url": "direct"}},
)
cache: dict[str, xai_auth_proxy.AuthFileSignature] = {}
xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
client.payloads[name]["proxy_url"] = "http://same-size-overwrite.invalid:1"
second = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
self.assertEqual(second.cached, 1)
self.assertEqual(second.downloaded, 2)
self.assertEqual(second.patched, 1)
self.assertEqual(len(client.patches), 2)
def test_provider_is_rechecked_immediately_before_patch(self) -> None:
name = "xai-race.json"
class RaceClient(FakeManagementClient):
def download_auth_file(self, auth_name: str) -> dict[str, Any]:
payload = super().download_auth_file(auth_name)
if len(self.downloads) == 1:
self.payloads[auth_name] = {"type": "codex", "proxy_url": "direct"}
return payload
client = RaceClient(
[{"name": name, "type": "xai", "modtime": "one", "size": 10}],
{name: {"type": "xai", "proxy_url": "direct"}},
)
result = xai_auth_proxy.enforce_management_once(client, PROXIES, {})
self.assertEqual(result.downloaded, 2)
self.assertEqual(result.patched, 0)
self.assertEqual(client.patches, [])
def test_downloaded_non_xai_is_never_patched_even_if_list_is_stale(self) -> None:
name = "shared-name.json"
client = FakeManagementClient(
[{"name": name, "type": "xai", "modtime": "one", "size": 10}],
{name: {"type": "codex", "proxy_url": "direct"}},
)
cache: dict[str, xai_auth_proxy.AuthFileSignature] = {}
result = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
self.assertEqual(result.downloaded, 1)
self.assertEqual(result.patched, 0)
self.assertEqual(client.patches, [])
self.assertNotIn(name, cache)
# The list signature can remain unchanged while its formerly stale
# download view converges to xAI. It must be inspected again.
client.payloads[name] = {"type": "xai", "proxy_url": "direct"}
converged = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
self.assertEqual(converged.downloaded, 2)
self.assertEqual(converged.patched, 1)
def test_failed_download_is_not_cached_and_is_retried(self) -> None:
name = "xai-retry.json"
client = FakeManagementClient(
[{"name": name, "type": "xai", "modtime": "one", "size": 10}],
{name: {"type": "xai"}},
)
client.download_error = RuntimeError("temporary failure")
cache: dict[str, xai_auth_proxy.AuthFileSignature] = {}
first = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
second = xai_auth_proxy.enforce_management_once(client, PROXIES, cache)
self.assertEqual(len(first.errors), 1)
self.assertNotEqual(second.patched, 0)
self.assertEqual(client.downloads, [name, name, name])
def test_watch_loop_has_injectable_sleep_and_cycle_limit(self) -> None:
name = "xai-watch.json"
proxy = xai_auth_proxy.stable_proxy_for_auth(name, PROXIES)
client = FakeManagementClient(
[{"name": name, "type": "xai", "modtime": "one", "size": 10}],
{name: {"type": "xai", "proxy_url": proxy}},
)
sleeps: list[float] = []
cycles = xai_auth_proxy.watch_management(
client,
PROXIES,
interval=0.25,
sleep=sleeps.append,
max_cycles=2,
)
self.assertEqual(cycles, 2)
self.assertEqual(sleeps, [0.25])
self.assertEqual(client.downloads, [name, name])
def test_watch_loop_keeps_retrying_when_management_is_initially_unavailable(self) -> None:
class UnavailableClient:
def list_auth_files(self) -> list[dict[str, Any]]:
raise RuntimeError("management unavailable")
errors: list[str] = []
sleeps: list[float] = []
cycles = xai_auth_proxy.watch_management(
UnavailableClient(),
PROXIES,
interval=0.1,
sleep=sleeps.append,
on_error=errors.append,
max_cycles=3,
)
self.assertEqual(cycles, 3)
self.assertEqual(sleeps, [0.1, 0.1])
self.assertEqual(errors, ["management unavailable"] * 3)
class CliTests(unittest.TestCase):
def test_once_cli_uses_proxy_environment(self) -> None:
with tempfile.TemporaryDirectory() as directory:
path = Path(directory) / "xai-cli.json"
path.write_text('{"type":"xai"}', encoding="utf-8")
stderr = io.StringIO()
with redirect_stderr(stderr):
status_code = xai_auth_proxy.main(
["once", "--auth-dir", directory],
{"XAI_PROXY_URLS": "\n".join(PROXIES)},
)
self.assertEqual(status_code, 0)
self.assertEqual(
json.loads(path.read_text(encoding="utf-8"))["proxy_url"],
xai_auth_proxy.stable_proxy_for_auth(path.name, PROXIES),
)
self.assertIn("updated=1", stderr.getvalue())
if __name__ == "__main__":
unittest.main()