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