feat: connect chat to managed model providers
This commit is contained in:
@@ -0,0 +1,360 @@
|
||||
#!/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())
|
||||
Reference in New Issue
Block a user