361 lines
13 KiB
Python
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": parsed.netloc,
|
|
"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())
|