from __future__ import annotations import importlib.util import json import os import shutil import stat import tempfile import unittest from pathlib import Path from unittest.mock import patch import xai_auth_proxy ROOT = Path(__file__).resolve().parent PROXIES = ("socks5h://127.0.0.1:1080", "socks5h://127.0.0.1:1081") def load_objectstore_sync(): required = { "OBJECTSTORE_ENDPOINT": "https://objectstore.invalid", "OBJECTSTORE_ACCESS_KEY": "test-access", "OBJECTSTORE_SECRET_KEY": "test-secret", "OBJECTSTORE_BUCKET": "test-bucket", "OBJECTSTORE_ROOT": "/tmp/objectstore-test-root", "XAI_PROXY_URLS": ",".join(PROXIES), } spec = importlib.util.spec_from_file_location( "objectstore_sync_under_test", ROOT / "objectstore_sync.py", ) assert spec is not None and spec.loader is not None module = importlib.util.module_from_spec(spec) with patch.dict(os.environ, required): spec.loader.exec_module(module) return module class ObjectstoreProxyEnforcementTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.module = load_objectstore_sync() def test_downloaded_xai_is_assigned_before_install(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) source = root / "remote-xai.json" source.write_text( '{"type":"xai","proxy_url":"direct","access_token":"A"}', encoding="utf-8", ) installed_root = root / "installed" uploads: list[str] = [] def fake_run_mc(args: list[str], check: bool = True): self.assertEqual(args[0], "cp") shutil.copyfile(source, args[2]) source.write_text( '{"type":"xai","proxy_url":"direct","access_token":"B"}', encoding="utf-8", ) original_root = self.module.ROOT original_run_mc = self.module.run_mc original_upload_file = self.module.upload_file try: self.module.ROOT = installed_root self.module.run_mc = fake_run_mc self.module.upload_file = uploads.append rel = "auths/xai-new.json" self.module.download_file( rel, { "remote": "objectstore/auths/xai-new.json", "last_modified": 1.0, }, ) finally: self.module.ROOT = original_root self.module.run_mc = original_run_mc self.module.upload_file = original_upload_file destination = installed_root / rel payload = json.loads(destination.read_text(encoding="utf-8")) self.assertEqual( payload["proxy_url"], xai_auth_proxy.stable_proxy_for_auth(destination.name, PROXIES), ) self.assertEqual(stat.S_IMODE(destination.stat().st_mode), 0o600) self.assertEqual(uploads, []) self.assertEqual(json.loads(source.read_text(encoding="utf-8"))["access_token"], "B") def test_non_xai_download_is_not_rewritten(self) -> None: with tempfile.TemporaryDirectory() as directory: path = Path(directory) / "codex.json" original = b'{ "type": "codex", "proxy_url": "provider-specific" }\n' path.write_bytes(original) changed = self.module._enforce_downloaded_xai_proxy(path, path.name) self.assertFalse(changed) self.assertEqual(path.read_bytes(), original) def test_remote_token_update_wins_on_followup_sync(self) -> None: with tempfile.TemporaryDirectory() as directory: root = Path(directory) installed_root = root / "installed" source = root / "remote-xai.json" rel = "auths/xai-race.json" remote_name = "REMOTE_XAI" source.write_text( '{"type":"xai","proxy_url":"direct","access_token":"A"}', encoding="utf-8", ) etag_a = self.module.file_md5(source) saved = { "ROOT": self.module.ROOT, "STATE_PATH": self.module.STATE_PATH, "run_mc": self.module.run_mc, "ensure_alias": self.module.ensure_alias, "enable_codex_auth_websockets": self.module.enable_codex_auth_websockets, "_fetch_invalid_auth_names": self.module._fetch_invalid_auth_names, "remote_inventory": self.module.remote_inventory, } try: self.module.ROOT = installed_root self.module.STATE_PATH = installed_root / ".objectstore-sync-state.json" def initial_download(args: list[str], check: bool = True): self.assertEqual(args[:2], ["cp", remote_name]) shutil.copyfile(source, args[2]) source.write_text( '{"type":"xai","proxy_url":"direct","access_token":"B"}', encoding="utf-8", ) self.module.run_mc = initial_download changed = self.module.download_file( rel, { "remote": remote_name, "etag": etag_a, "last_modified": 1.0, }, ) self.assertTrue(changed) derived: dict[str, dict[str, str]] = {} self.module.remember_derived_file(derived, rel, etag_a) self.module.save_known_rels(self.module.local_inventory(), derived) etag_b = self.module.file_md5(source) uploads: list[str] = [] def sync_mc(args: list[str], check: bool = True): self.assertEqual(args[0], "cp") if args[1] == remote_name: shutil.copyfile(source, args[2]) else: uploads.append(args[1]) shutil.copyfile(args[1], source) self.module.run_mc = sync_mc self.module.ensure_alias = lambda: None self.module.enable_codex_auth_websockets = lambda path: None self.module._fetch_invalid_auth_names = lambda: set() self.module.remote_inventory = lambda: { rel: { "remote": remote_name, "etag": etag_b, "last_modified": 2.0, } } self.module.sync() self.assertEqual(uploads, []) installed_path = installed_root / rel refreshed = json.loads(installed_path.read_text(encoding="utf-8")) refreshed["access_token"] = "C" installed_path.write_text(json.dumps(refreshed), encoding="utf-8") os.utime(installed_path, (3.0, 3.0)) self.module.sync() finally: for name, value in saved.items(): setattr(self.module, name, value) self.assertEqual(uploads, [str(installed_root / rel)]) self.assertEqual(json.loads(source.read_text(encoding="utf-8"))["access_token"], "C") installed = json.loads((installed_root / rel).read_text(encoding="utf-8")) self.assertEqual(installed["access_token"], "C") self.assertEqual( installed["proxy_url"], xai_auth_proxy.stable_proxy_for_auth(Path(rel).name, PROXIES), ) if __name__ == "__main__": unittest.main()