Files
CASAN/packages/casan-harness/scripts/bash/model-connections.py
T
2026-07-11 12:18:59 +09:00

361 lines
13 KiB
Python

#!/usr/bin/env python3
"""Tenant-scoped model connections for CASAN Chat.
Secrets are accepted only through stdin, encrypted with tenant-crypt.sh, and
never returned by public list/refresh operations. ``runtime-env`` is intended
only for the backend child process that launches a governed chat turn.
"""
import argparse
import json
import os
import subprocess
import sys
import tempfile
import urllib.error
import urllib.request
from datetime import datetime, timezone
from typing import Optional
from urllib.parse import urlparse
def project_root() -> str:
current = os.path.abspath(os.path.dirname(__file__))
while current != os.path.dirname(current):
if os.path.isdir(os.path.join(current, ".specify")):
return current
current = os.path.dirname(current)
raise SystemExit("MODEL_CONNECTION_ROOT_NOT_FOUND")
ROOT = project_root()
BIN = os.path.join(ROOT, "packages", "casan-harness", "scripts", "bash")
TENANT_STORE = os.path.join(BIN, "tenant-store.sh")
TENANT_CRYPT = os.path.join(BIN, "tenant-crypt.sh")
PROVIDERS = {
"openai": {
"label": "OpenAI / Codex API",
"kind": "cloud",
"policy_provider": "cloud-openai",
"default_endpoint": "https://api.openai.com/v1",
"model_prefix": "openai",
"requires_key": True,
},
"anthropic": {
"label": "Anthropic / Claude API",
"kind": "cloud",
"policy_provider": "cloud-anthropic",
"default_endpoint": "https://api.anthropic.com/v1",
"model_prefix": "anthropic",
"requires_key": True,
},
"omniroute": {
"label": "OmniRoute Gateway",
"kind": "gateway",
"policy_provider": "omniroute",
"default_endpoint": "http://host.docker.internal:20128/v1",
"model_prefix": "openai-compatible",
"requires_key": False,
},
"ollama": {
"label": "Ollama Local",
"kind": "local",
"policy_provider": "local",
"default_endpoint": "http://host.docker.internal:11434",
"model_prefix": "ollama",
"requires_key": False,
},
}
def now() -> str:
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
def fail(message: str, code: int = 2) -> None:
print(json.dumps({"success": False, "reason": message}), file=sys.stderr)
raise SystemExit(code)
def run_bash(script: str, args: list[str]) -> subprocess.CompletedProcess[str]:
return subprocess.run(["bash", script, *args], cwd=ROOT, capture_output=True, text=True)
def store_path() -> str:
result = run_bash(TENANT_STORE, ["resolve", "model-connections/connections.json.enc"])
if result.returncode != 0:
fail((result.stderr or result.stdout or "TENANT_STORE_DENIED").strip(), 3)
return result.stdout.strip()
def load_store() -> dict:
path = store_path()
if not os.path.isfile(path):
return {"version": 1, "connections": {}}
with tempfile.TemporaryDirectory() as tmp:
plain = os.path.join(tmp, "connections.json")
result = run_bash(TENANT_CRYPT, ["decrypt", path, plain])
if result.returncode != 0:
fail("MODEL_CONNECTION_DECRYPT_FAILED", 3)
try:
with open(plain, encoding="utf-8") as handle:
payload = json.load(handle)
except (OSError, ValueError):
fail("MODEL_CONNECTION_STORE_INVALID", 3)
return payload if isinstance(payload, dict) else {"version": 1, "connections": {}}
def save_store(payload: dict) -> None:
destination = store_path()
os.makedirs(os.path.dirname(destination), exist_ok=True)
with tempfile.TemporaryDirectory(dir=os.path.dirname(destination)) as tmp:
plain = os.path.join(tmp, "connections.json")
encrypted = os.path.join(tmp, "connections.json.enc")
with open(plain, "w", encoding="utf-8") as handle:
json.dump(payload, handle, ensure_ascii=False, sort_keys=True)
os.chmod(plain, 0o600)
result = run_bash(TENANT_CRYPT, ["encrypt", plain, encrypted])
if result.returncode != 0:
fail("MODEL_CONNECTION_ENCRYPT_FAILED", 3)
os.chmod(encrypted, 0o600)
os.replace(encrypted, destination)
def provider_config(provider_id: str) -> dict:
config = PROVIDERS.get(provider_id)
if not config:
fail("unknown_provider")
return config
def allowed_local_hosts() -> set[str]:
configured = os.environ.get("CASAN_MODEL_CONNECTION_ALLOWED_HOSTS", "")
return {"127.0.0.1", "localhost", "host.docker.internal"} | {
item.strip().lower() for item in configured.split(",") if item.strip()
}
def validate_endpoint(provider_id: str, endpoint: str) -> str:
config = provider_config(provider_id)
value = (endpoint or config["default_endpoint"]).strip().rstrip("/")
parsed = urlparse(value)
if parsed.scheme not in {"http", "https"} or not parsed.hostname or parsed.username or parsed.password:
fail("endpoint_invalid")
if parsed.query or parsed.fragment:
fail("endpoint_invalid")
host = parsed.hostname.lower()
if provider_id == "openai" and value != "https://api.openai.com/v1":
fail("endpoint_not_allowed")
if provider_id == "anthropic" and value != "https://api.anthropic.com/v1":
fail("endpoint_not_allowed")
if provider_id in {"omniroute", "ollama"}:
if host not in allowed_local_hosts():
fail("endpoint_not_allowed")
if parsed.scheme == "http" and host not in allowed_local_hosts():
fail("insecure_endpoint_not_allowed")
allowed_paths = {"", "/v1"} if provider_id == "omniroute" else {""}
if parsed.path.rstrip("/") not in allowed_paths:
fail("endpoint_path_not_allowed")
return value
def request_json(url: str, headers: dict[str, str]) -> dict:
request = urllib.request.Request(url, headers=headers)
try:
with urllib.request.urlopen(request, timeout=12) as response:
return json.loads(response.read().decode("utf-8"))
except urllib.error.HTTPError as exc:
fail(f"provider_http_{exc.code}")
except (urllib.error.URLError, TimeoutError, ValueError):
fail("provider_unreachable")
return {}
def discover(provider_id: str, endpoint: str, api_key: str) -> list[str]:
if provider_id == "ollama":
payload = request_json(f"{endpoint}/api/tags", {})
rows = payload.get("models", [])
ids = [row.get("name", "") for row in rows if isinstance(row, dict)]
else:
headers = {"Accept": "application/json"}
if provider_id == "anthropic":
headers.update({"x-api-key": api_key, "anthropic-version": "2023-06-01"})
else:
headers["Authorization"] = f"Bearer {api_key}"
payload = request_json(f"{endpoint}/models", headers)
rows = payload.get("data", payload.get("models", []))
excluded_types = {"image", "video", "audio", "embedding", "rerank"}
ids = [
row.get("id", row.get("name", ""))
for row in rows
if isinstance(row, dict) and str(row.get("type", "")).lower() not in excluded_types
]
return sorted({str(model_id).strip() for model_id in ids if str(model_id).strip()})
def public_connection(provider_id: str, saved: Optional[dict]) -> dict:
config = provider_config(provider_id)
row = saved or {}
return {
"id": provider_id,
"label": config["label"],
"kind": config["kind"],
"connected": bool(saved),
"endpoint": row.get("endpoint", config["default_endpoint"]),
"models": row.get("models", []),
"defaultModel": row.get("default_model", ""),
"lastCheckedAt": row.get("last_checked_at"),
"status": row.get("status", "not_connected"),
"requiresKey": config["requires_key"],
}
def list_connections() -> int:
connections = load_store().get("connections", {})
print(json.dumps({
"success": True,
"connections": [public_connection(provider_id, connections.get(provider_id)) for provider_id in PROVIDERS],
}, ensure_ascii=False))
return 0
def stdin_payload() -> dict:
try:
payload = json.load(sys.stdin)
except ValueError:
fail("invalid_request")
return payload if isinstance(payload, dict) else {}
def connect(provider_id: str) -> int:
config = provider_config(provider_id)
body = stdin_payload()
api_key = str(body.get("apiKey", "")).strip()
if config["requires_key"] and not api_key:
fail("api_key_required")
if provider_id == "omniroute" and not api_key:
api_key = os.environ.get("CASAN_OPENAI_COMPATIBLE_API_KEY", "local-omniroute")
endpoint = validate_endpoint(provider_id, str(body.get("endpoint", "")))
models = discover(provider_id, endpoint, api_key)
if not models:
fail("provider_returned_no_models")
payload = load_store()
payload.setdefault("connections", {})[provider_id] = {
"endpoint": endpoint,
"api_key": api_key,
"models": models,
"default_model": str(body.get("defaultModel", "")) if body.get("defaultModel") in models else models[0],
"last_checked_at": now(),
"status": "connected",
}
save_store(payload)
print(json.dumps({"success": True, "connection": public_connection(provider_id, payload["connections"][provider_id])}, ensure_ascii=False))
return 0
def refresh(provider_id: str) -> int:
provider_config(provider_id)
payload = load_store()
saved = payload.get("connections", {}).get(provider_id)
if not saved:
fail("provider_not_connected")
models = discover(provider_id, validate_endpoint(provider_id, saved.get("endpoint", "")), saved.get("api_key", ""))
if not models:
fail("provider_returned_no_models")
saved["models"] = models
saved["last_checked_at"] = now()
saved["status"] = "connected"
if saved.get("default_model") not in models:
saved["default_model"] = models[0]
save_store(payload)
print(json.dumps({"success": True, "connection": public_connection(provider_id, saved)}, ensure_ascii=False))
return 0
def set_default(provider_id: str, model_id: str) -> int:
payload = load_store()
saved = payload.get("connections", {}).get(provider_id)
if not saved:
fail("provider_not_connected")
if model_id not in saved.get("models", []):
fail("model_not_discovered")
saved["default_model"] = model_id
save_store(payload)
print(json.dumps({"success": True, "connection": public_connection(provider_id, saved)}, ensure_ascii=False))
return 0
def disconnect(provider_id: str) -> int:
provider_config(provider_id)
payload = load_store()
payload.get("connections", {}).pop(provider_id, None)
save_store(payload)
print(json.dumps({"success": True, "connection": public_connection(provider_id, None)}, ensure_ascii=False))
return 0
def runtime_env(provider_id: str, model_id: str) -> int:
config = provider_config(provider_id)
saved = load_store().get("connections", {}).get(provider_id)
if not saved:
fail("provider_not_connected")
selected = model_id or saved.get("default_model", "")
if selected not in saved.get("models", []):
fail("model_not_discovered")
endpoint = validate_endpoint(provider_id, saved.get("endpoint", ""))
env = {
"CASAN_CHAT_MODEL_MODE": "model",
"CASAN_CHAT_MODEL_PROVIDER": config["policy_provider"],
"CASAN_CHAT_SELECTED_MODEL": f"{config['model_prefix']}:{selected}",
}
if provider_id == "openai":
env["OPENAI_API_KEY"] = saved.get("api_key", "")
elif provider_id == "anthropic":
env["ANTHROPIC_API_KEY"] = saved.get("api_key", "")
elif provider_id == "omniroute":
parsed = urlparse(endpoint)
env.update({
"CASAN_OPENAI_COMPATIBLE_API_KEY": saved.get("api_key", ""),
"CASAN_OPENAI_COMPATIBLE_BASE_URL": endpoint,
"CASAN_OPENAI_COMPATIBLE_ALLOWED_HOSTS": parsed.hostname or "",
})
else:
parsed = urlparse(endpoint)
env.update({
"CASAN_OLLAMA_HOST": parsed.netloc,
"OLLAMA_HOST": endpoint,
"CASAN_ALLOW_DOCKER_HOST_OLLAMA": "1" if parsed.hostname == "host.docker.internal" else "0",
})
print(json.dumps({"success": True, "env": env}, ensure_ascii=False))
return 0
def main() -> int:
parser = argparse.ArgumentParser()
sub = parser.add_subparsers(dest="command", required=True)
sub.add_parser("list")
for name in ("connect", "refresh", "disconnect"):
command = sub.add_parser(name)
command.add_argument("--provider", required=True)
default = sub.add_parser("set-default")
default.add_argument("--provider", required=True)
default.add_argument("--model", required=True)
runtime = sub.add_parser("runtime-env")
runtime.add_argument("--provider", required=True)
runtime.add_argument("--model", default="")
args = parser.parse_args()
if args.command == "list":
return list_connections()
if args.command == "connect":
return connect(args.provider)
if args.command == "refresh":
return refresh(args.provider)
if args.command == "set-default":
return set_default(args.provider, args.model)
if args.command == "disconnect":
return disconnect(args.provider)
return runtime_env(args.provider, args.model)
if __name__ == "__main__":
raise SystemExit(main())