daili / test_objectstore_xai_proxy.py
pjpjq's picture
fix(proxy): 强制 xAI 认证使用多出口代理
bc3a943
Raw
History Blame Contribute Delete
7.92 kB
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()