Source code for kumiho.discovery

"""Helpers for bootstrapping a Client via the control-plane discovery endpoint."""

from __future__ import annotations

import base64
import binascii
import hashlib
import hmac
import ipaddress
import json
import os
import platform
import secrets
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import TYPE_CHECKING, Any, Dict, Iterable, Optional, Sequence, Tuple, Type
from urllib.parse import urlparse

import requests

from ._token_loader import load_bearer_token, load_firebase_token
if TYPE_CHECKING:
    from .client import _Client as ClientType
else:
    ClientType = Any

Client: Optional[Type[Any]] = None

DEFAULT_CONTROL_PLANE_URL = os.getenv("KUMIHO_CONTROL_PLANE_URL") or "https://control.kumiho.cloud"
DEFAULT_CACHE_PATH = Path(
    os.getenv("KUMIHO_DISCOVERY_CACHE_FILE")
    or (Path.home() / ".kumiho" / "discovery-cache.json")
)
_DEFAULT_TIMEOUT = float(os.getenv("KUMIHO_DISCOVERY_TIMEOUT_SECONDS", "10"))
_DEFAULT_CACHE_KEY = "__default__"
_LOCAL_CE_ENDPOINT_ENV = "KUMIHO_LOCAL_SERVER_ENDPOINT"
_LOCAL_CE_PORT_ENV = "KUMIHO_LOCAL_SERVER_PORT"
_LOCAL_CE_TIMEOUT_ENV = "KUMIHO_LOCAL_DISCOVERY_TIMEOUT_SECONDS"
# Self-hosted CE default loopback port. Must match the server's CE default
# (kumiho-server config/apply_deployment_defaults) and the installer scripts.
_DEFAULT_LOCAL_CE_PORT = 9190
_DEFAULT_LOCAL_CE_TARGET = f"127.0.0.1:{_DEFAULT_LOCAL_CE_PORT}"


class DiscoveryError(RuntimeError):
    """Raised when the discovery endpoint cannot be reached or returns an error."""


@dataclass(frozen=True)
class RegionRouting:
    region_code: str
    server_url: str
    grpc_authority: Optional[str] = None

    @classmethod
    def from_dict(cls, payload: Dict[str, Any]) -> "RegionRouting":
        return cls(
            region_code=payload["region_code"],
            server_url=payload["server_url"],
            grpc_authority=payload.get("grpc_authority"),
        )

    def to_dict(self) -> Dict[str, Any]:
        data: Dict[str, Any] = {
            "region_code": self.region_code,
            "server_url": self.server_url,
        }
        if self.grpc_authority:
            data["grpc_authority"] = self.grpc_authority
        return data


@dataclass(frozen=True)
class CacheControl:
    issued_at: datetime
    refresh_at: datetime
    expires_at: datetime
    expires_in_seconds: int
    refresh_after_seconds: int

    @classmethod
    def from_dict(cls, payload: Dict[str, Any]) -> "CacheControl":
        issued_at = _parse_iso8601(payload.get("issued_at"))
        refresh_at = _parse_iso8601(payload.get("refresh_at"))
        expires_at = _parse_iso8601(payload.get("expires_at"))
        return cls(
            issued_at=issued_at,
            refresh_at=refresh_at,
            expires_at=expires_at,
            expires_in_seconds=int(payload.get("expires_in_seconds", 0)),
            refresh_after_seconds=int(payload.get("refresh_after_seconds", 0)),
        )

    def to_dict(self) -> Dict[str, Any]:
        return {
            "issued_at": self.issued_at.isoformat(),
            "refresh_at": self.refresh_at.isoformat(),
            "expires_at": self.expires_at.isoformat(),
            "expires_in_seconds": self.expires_in_seconds,
            "refresh_after_seconds": self.refresh_after_seconds,
        }

    def is_expired(self, *, now: Optional[datetime] = None) -> bool:
        moment = now or datetime.now(timezone.utc)
        return moment >= self.expires_at

    def should_refresh(self, *, now: Optional[datetime] = None) -> bool:
        moment = now or datetime.now(timezone.utc)
        return moment >= self.refresh_at


@dataclass(frozen=True)
class DiscoveryRecord:
    tenant_id: str
    tenant_name: Optional[str]
    roles: Sequence[str]
    guardrails: Optional[Dict[str, Any]]
    region: RegionRouting
    cache_control: CacheControl

    @classmethod
    def from_dict(cls, payload: Dict[str, Any]) -> "DiscoveryRecord":
        cache_section = payload.get("cache_control")
        if not cache_section:
            raise DiscoveryError("Discovery payload is missing cache_control metadata")
        region_section = payload.get("region")
        if not region_section:
            raise DiscoveryError("Discovery payload is missing region metadata")
        return cls(
            tenant_id=payload["tenant_id"],
            tenant_name=payload.get("tenant_name"),
            roles=list(payload.get("roles", [])),
            guardrails=payload.get("guardrails"),
            region=RegionRouting.from_dict(region_section),
            cache_control=CacheControl.from_dict(cache_section),
        )

    def to_dict(self) -> Dict[str, Any]:
        return {
            "tenant_id": self.tenant_id,
            "tenant_name": self.tenant_name,
            "roles": list(self.roles),
            "guardrails": self.guardrails,
            "region": self.region.to_dict(),
            "cache_control": self.cache_control.to_dict(),
        }


def _get_machine_id() -> str:
    """Get a stable machine identifier for deriving encryption keys.
    
    This provides defense-in-depth by making cache files non-portable
    between machines. Falls back to a random ID stored in config dir.
    """
    # Try platform-specific methods
    try:
        if platform.system() == "Linux":
            # Linux: use machine-id
            for path in ["/etc/machine-id", "/var/lib/dbus/machine-id"]:
                if os.path.exists(path):
                    with open(path, "r") as f:
                        return f.read().strip()
        elif platform.system() == "Darwin":
            # macOS: use hardware UUID
            import subprocess
            result = subprocess.run(
                ["ioreg", "-rd1", "-c", "IOPlatformExpertDevice"],
                capture_output=True,
                text=True,
                timeout=5,
            )
            for line in result.stdout.split("\n"):
                if "IOPlatformUUID" in line:
                    return line.split('"')[-2]
        elif platform.system() == "Windows":
            # Windows: use machine GUID
            import winreg
            key = winreg.OpenKey(
                winreg.HKEY_LOCAL_MACHINE,
                r"SOFTWARE\Microsoft\Cryptography",
            )
            value, _ = winreg.QueryValueEx(key, "MachineGuid")
            return str(value)
    except Exception:
        pass  # Fall through to file-based ID
    
    # Fallback: use a randomly generated ID stored in config dir
    from ._token_loader import _config_dir
    id_file = _config_dir() / ".machine_id"
    if id_file.exists():
        return id_file.read_text(encoding="utf-8").strip()
    
    # Generate and store new ID
    new_id = str(uuid.uuid4())
    id_file.parent.mkdir(parents=True, exist_ok=True)
    id_file.write_text(new_id, encoding="utf-8")
    return new_id


def _derive_cache_key() -> bytes:
    """Derive an encryption key from machine ID + user context."""
    machine_id = _get_machine_id()
    try:
        login = os.getlogin()
    except OSError:
        login = ""
    uid = str(os.getuid()) if hasattr(os, "getuid") else ""
    user_context = f"{login}{uid}"
    key_material = f"kumiho-discovery-cache-v1:{machine_id}:{user_context}"
    return hashlib.sha256(key_material.encode()).digest()


def _encrypt_cache_data(plaintext: str) -> str:
    """Encrypt cache data using XOR cipher with HMAC authentication.
    
    This provides defense-in-depth, not cryptographic security against
    determined attackers. The goal is to prevent casual inspection and
    make cache files non-portable between machines.
    """
    key = _derive_cache_key()
    
    # Generate random IV
    iv = secrets.token_bytes(16)
    
    # XOR encryption (simple but effective for this use case)
    plaintext_bytes = plaintext.encode("utf-8")
    key_stream = hashlib.sha256(key + iv).digest()
    
    # Extend key stream for longer plaintexts
    while len(key_stream) < len(plaintext_bytes):
        key_stream += hashlib.sha256(key + key_stream[-32:]).digest()
    
    ciphertext = bytes(p ^ k for p, k in zip(plaintext_bytes, key_stream))
    
    # Add HMAC for integrity
    mac = hmac.new(key, iv + ciphertext, hashlib.sha256).digest()[:16]
    
    # Format: base64(iv + ciphertext + mac)
    encrypted = base64.b64encode(iv + ciphertext + mac).decode("ascii")
    return f"enc:v1:{encrypted}"


def _decrypt_cache_data(encrypted: str) -> Optional[str]:
    """Decrypt cache data. Returns None if decryption fails."""
    if not encrypted.startswith("enc:v1:"):
        # Unencrypted legacy format - migrate on next write
        return encrypted
    
    try:
        key = _derive_cache_key()
        raw = base64.b64decode(encrypted[7:])  # Skip "enc:v1:" prefix
        
        if len(raw) < 32:  # iv(16) + mac(16) minimum
            return None
        
        iv = raw[:16]
        mac = raw[-16:]
        ciphertext = raw[16:-16]
        
        # Verify HMAC
        expected_mac = hmac.new(key, iv + ciphertext, hashlib.sha256).digest()[:16]
        if not hmac.compare_digest(mac, expected_mac):
            return None  # Integrity check failed
        
        # Decrypt
        key_stream = hashlib.sha256(key + iv).digest()
        while len(key_stream) < len(ciphertext):
            key_stream += hashlib.sha256(key + key_stream[-32:]).digest()
        
        plaintext = bytes(c ^ k for c, k in zip(ciphertext, key_stream))
        return plaintext.decode("utf-8")
    except Exception:
        return None


class DiscoveryCache:
    """Encrypted JSON file cache keyed by tenant hint.
    
    Cache data is encrypted at rest using a machine-specific key,
    providing defense-in-depth protection for tenant metadata.
    """

    def __init__(self, path: Optional[Path] = None, *, encrypt: bool = True) -> None:
        self.path = path or DEFAULT_CACHE_PATH
        self._encrypt = encrypt

    def load(self, cache_key: str) -> Optional[DiscoveryRecord]:
        payload = self._read_all().get(cache_key)
        if not payload:
            return None
        try:
            return DiscoveryRecord.from_dict(payload)
        except DiscoveryError:
            return None

    def store(self, cache_key: str, record: DiscoveryRecord) -> None:
        data = self._read_all()
        data[cache_key] = record.to_dict()
        self.path.parent.mkdir(parents=True, exist_ok=True)
        tmp_path = self.path.with_suffix(".tmp")
        
        # Serialize to JSON
        json_content = json.dumps(data, indent=2)
        
        # Encrypt if enabled
        if self._encrypt:
            content_to_write = _encrypt_cache_data(json_content)
        else:
            content_to_write = json_content
        
        with tmp_path.open("w", encoding="utf-8") as handle:
            handle.write(content_to_write)
        
        # Retry replacement to handle Windows file locking
        import time
        max_retries = 5
        for i in range(max_retries):
            try:
                tmp_path.replace(self.path)
                return
            except PermissionError:
                if i == max_retries - 1:
                    raise
                time.sleep(0.1)
            except OSError:
                if i == max_retries - 1:
                    raise
                time.sleep(0.1)

    def _read_all(self) -> Dict[str, Any]:
        if not self.path.exists():
            return {}
        try:
            with self.path.open("r", encoding="utf-8") as handle:
                content = handle.read()
            
            # Try to decrypt if encrypted
            decrypted = _decrypt_cache_data(content)
            if decrypted is None:
                # Decryption failed - cache may be from different machine
                return {}
            
            return json.loads(decrypted)
        except (json.JSONDecodeError, OSError):
            return {}


class DiscoveryManager:
    """Coordinates cache usage and remote discovery calls."""

    def __init__(
        self,
        *,
        control_plane_url: Optional[str] = None,
        cache_path: Optional[Path] = None,
        timeout: Optional[float] = None,
    ) -> None:
        self.base_url = control_plane_url or DEFAULT_CONTROL_PLANE_URL
        self.cache = DiscoveryCache(cache_path)
        self.timeout = timeout or _DEFAULT_TIMEOUT

    def resolve(
        self,
        *,
        id_token: str,
        tenant_hint: Optional[str] = None,
        force_refresh: bool = False,
    ) -> DiscoveryRecord:
        cache_key = tenant_hint or _DEFAULT_CACHE_KEY

        def fetch_fresh() -> DiscoveryRecord:
            last_error: Optional[DiscoveryError] = None
            for token in _discovery_token_candidates(id_token):
                try:
                    fresh = self._fetch_remote(id_token=token, tenant_hint=tenant_hint)
                    self.cache.store(cache_key, fresh)
                    return fresh
                except DiscoveryError as exc:
                    last_error = exc
                    continue

            if last_error:
                raise last_error
            raise DiscoveryError("Discovery failed without a usable bearer token")

        if not force_refresh:
            cached = self.cache.load(cache_key)
            if cached and not cached.cache_control.is_expired():
                if cached.cache_control.should_refresh():
                    try:
                        return fetch_fresh()
                    except DiscoveryError:
                        if not cached.cache_control.is_expired():
                            return cached
                        raise
                return cached

        return fetch_fresh()

    def _fetch_remote(self, *, id_token: str, tenant_hint: Optional[str]) -> DiscoveryRecord:
        from kumiho import __version__ as _sdk_version

        url = _build_discovery_url(self.base_url)
        headers = {
            "Authorization": f"Bearer {id_token}",
            "Content-Type": "application/json",
            "User-Agent": f"kumiho-python/{_sdk_version}",
        }
        payload: Dict[str, Any] = {}
        if tenant_hint:
            payload["tenant_hint"] = tenant_hint

        response = requests.post(url, json=payload, headers=headers, timeout=self.timeout)
        if response.status_code >= 400:
            raise DiscoveryError(
                f"Discovery endpoint returned {response.status_code}: {response.text[:200]}"
            )
        try:
            body = response.json()
        except ValueError as exc:
            raise DiscoveryError("Discovery endpoint returned invalid JSON") from exc
        return DiscoveryRecord.from_dict(body)


def client_from_discovery(
    *,
    id_token: Optional[str] = None,
    tenant_hint: Optional[str] = None,
    control_plane_url: Optional[str] = None,
    cache_path: Optional[str] = None,
    force_refresh: bool = False,
    default_metadata: Optional[Sequence[Tuple[str, str]]] = None,
) -> "ClientType":
    """Create a Client configured via the public discovery endpoint.

    The helper caches discovery payloads based on the tenant hint, respects the
    cache-control metadata emitted by the control plane, and refreshes the
    routing info once the `refresh_after_seconds` deadline passes.
    """

    token = id_token or load_bearer_token()
    if not token:
        raise DiscoveryError(
            "A bearer token is required. Set KUMIHO_AUTH_TOKEN or run kumiho-auth login."
        )

    manager = DiscoveryManager(
        control_plane_url=control_plane_url,
        cache_path=Path(cache_path) if cache_path else None,
    )
    record = manager.resolve(id_token=token, tenant_hint=tenant_hint, force_refresh=force_refresh)

    target = record.region.grpc_authority or record.region.server_url
    metadata: Iterable[Tuple[str, str]] = list(default_metadata or [])
    metadata = list(metadata)
    metadata.append(("x-tenant-id", record.tenant_id))

    client_cls = _get_client_class()
    return client_cls(target=target, auth_token=token, default_metadata=metadata)


[docs] def client_from_local_ce( *, default_metadata: Optional[Sequence[Tuple[str, str]]] = None, timeout: Optional[float] = None, ) -> Optional["ClientType"]: """Create a tokenless Client for a loopback self-hosted CE server, if present.""" target = resolve_local_ce_endpoint(timeout=timeout) if not target: return None client_cls = _get_client_class() return client_cls( target=target, auth_token=None, default_metadata=default_metadata, use_discovery=False, enable_auto_login=False, skip_auth_token_load=True, )
def resolve_local_ce_endpoint(*, timeout: Optional[float] = None) -> Optional[str]: """Return a loopback CE gRPC target when a local server advertises CE mode.""" probe_timeout = timeout if timeout is not None else _local_ce_timeout() for target in _local_ce_candidates(): if _probe_local_ce_candidate(target, timeout=probe_timeout): return target return None def _local_ce_timeout() -> float: raw = os.getenv(_LOCAL_CE_TIMEOUT_ENV) if not raw: return 0.5 try: return max(float(raw), 0.05) except ValueError: return 0.5 def _local_ce_candidates() -> Sequence[str]: endpoint = os.getenv(_LOCAL_CE_ENDPOINT_ENV) if endpoint and endpoint.strip(): return [_normalise_local_ce_target(endpoint)] port = os.getenv(_LOCAL_CE_PORT_ENV) if port and port.strip(): if not port.strip().isdigit(): raise DiscoveryError(f"{_LOCAL_CE_PORT_ENV} must be a numeric loopback port") return [f"127.0.0.1:{int(port.strip())}"] return [_DEFAULT_LOCAL_CE_TARGET] def _normalise_local_ce_target(raw: str) -> str: text = raw.strip() parsed = urlparse(text if "://" in text else f"//{text}") host = parsed.hostname if not host: raise DiscoveryError(f"{_LOCAL_CE_ENDPOINT_ENV} must include a loopback host") if not _is_loopback_host(host): raise DiscoveryError( f"{_LOCAL_CE_ENDPOINT_ENV} must point to localhost, 127.0.0.1, or ::1" ) port = parsed.port or _DEFAULT_LOCAL_CE_PORT if port <= 0 or port > 65535: raise DiscoveryError(f"{_LOCAL_CE_ENDPOINT_ENV} port must be between 1 and 65535") return f"{_format_host_for_target(host)}:{port}" def _is_loopback_host(host: str) -> bool: if host.lower() == "localhost": return True try: return ipaddress.ip_address(host).is_loopback except ValueError: return False def _format_host_for_target(host: str) -> str: if ":" in host and not host.startswith("["): return f"[{host}]" return host def _probe_local_ce_candidate(target: str, *, timeout: float) -> bool: url = f"http://{target}/api/_live" try: response = requests.get(url, timeout=timeout) except requests.RequestException: return False if response.status_code >= 400: return False try: body = response.json() except ValueError: return False return body.get("deployment_mode") == "self_hosted_ce" def _parse_iso8601(raw: Optional[str]) -> datetime: if not raw: raise DiscoveryError("Discovery payload missing required timestamp") text = raw.strip() if text.endswith("Z"): text = text[:-1] + "+00:00" value = datetime.fromisoformat(text) if value.tzinfo is None: value = value.replace(tzinfo=timezone.utc) return value.astimezone(timezone.utc) def _build_discovery_url(base_url: str) -> str: base = base_url.rstrip("/") if base.endswith("/api/discovery/tenant"): return base if base.endswith("/api/discovery"): return f"{base}/tenant" if base.endswith("/api"): return f"{base}/discovery/tenant" return f"{base}/api/discovery/tenant" def _discovery_token_candidates(candidate: str) -> list[str]: candidates = [candidate] if not _is_control_plane_token(candidate): return candidates firebase = load_firebase_token() if firebase and firebase not in candidates: candidates.append(firebase) return candidates def _decode_claims(token: str) -> Dict[str, Any]: parts = token.split(".") if len(parts) < 2: return {} payload = parts[1] padding = "=" * (-len(payload) % 4) try: data = base64.urlsafe_b64decode((payload + padding).encode("utf-8")) except (binascii.Error, ValueError): return {} try: obj = json.loads(data) except json.JSONDecodeError: return {} return obj if isinstance(obj, dict) else {} def _is_control_plane_token(token: str) -> bool: claims = _decode_claims(token) if not claims: return False if isinstance(claims.get("tenant_id"), str): return True iss = claims.get("iss") if isinstance(iss, str) and iss.startswith("https://control.kumiho.cloud"): return True aud = claims.get("aud") if isinstance(aud, str) and aud.startswith("kumiho-server"): return True return False def _get_client_class() -> Type[Any]: global Client if Client is None: from .client import _Client as RealClient Client = RealClient return Client