97 lines
3.3 KiB
Python
97 lines
3.3 KiB
Python
"""Centralized runtime configuration for Oracle embedding inference."""
|
|
|
|
import os
|
|
from collections.abc import Mapping
|
|
from dataclasses import dataclass
|
|
from pathlib import Path
|
|
from urllib.parse import urlsplit
|
|
|
|
ORACLE_BASE_URL_VARIABLE = "THOTH_ORACLE_BASE_URL"
|
|
EMBEDDING_MODEL_VARIABLE = "THOTH_EMBEDDING_MODEL"
|
|
EMBEDDING_DIMENSION_VARIABLE = "THOTH_EMBEDDING_DIMENSION"
|
|
|
|
|
|
@dataclass(frozen=True, slots=True)
|
|
class EmbeddingConfig:
|
|
"""Deployment values used to construct an Oracle embedding provider."""
|
|
|
|
oracle_base_url: str | None
|
|
embedding_model: str | None
|
|
embedding_dimension: int | None
|
|
|
|
def require_live_embedding(self) -> "EmbeddingConfig":
|
|
"""Validate values required before making a live inference request."""
|
|
|
|
if not self.oracle_base_url:
|
|
raise ValueError(f"{ORACLE_BASE_URL_VARIABLE} is required for embeddings")
|
|
parsed_url = urlsplit(self.oracle_base_url)
|
|
if parsed_url.scheme not in {"http", "https"} or not parsed_url.netloc:
|
|
raise ValueError(
|
|
f"{ORACLE_BASE_URL_VARIABLE} must be an explicit HTTP(S) URL "
|
|
"including its scheme"
|
|
)
|
|
if parsed_url.query or parsed_url.fragment:
|
|
raise ValueError(
|
|
f"{ORACLE_BASE_URL_VARIABLE} must not contain a query or fragment"
|
|
)
|
|
if not self.embedding_model:
|
|
raise ValueError(f"{EMBEDDING_MODEL_VARIABLE} is required for embeddings")
|
|
return self
|
|
|
|
|
|
def read_env_file(path: Path) -> dict[str, str]:
|
|
"""Read simple KEY=VALUE entries from a local dotenv file."""
|
|
|
|
if not path.exists():
|
|
return {}
|
|
|
|
values: dict[str, str] = {}
|
|
for line_number, original_line in enumerate(
|
|
path.read_text(encoding="utf-8").splitlines(), start=1
|
|
):
|
|
line = original_line.strip()
|
|
if not line or line.startswith("#"):
|
|
continue
|
|
if "=" not in line:
|
|
raise ValueError(f"invalid .env entry at {path}:{line_number}")
|
|
name, value = line.split("=", 1)
|
|
name = name.strip()
|
|
value = value.strip()
|
|
if value[:1] == value[-1:] and value[:1] in {'"', "'"}:
|
|
value = value[1:-1]
|
|
values[name] = value
|
|
return values
|
|
|
|
|
|
def load_embedding_config(
|
|
environ: Mapping[str, str] | None = None,
|
|
env_path: str | Path = ".env",
|
|
) -> EmbeddingConfig:
|
|
"""Load `.env`, then override it with operating-system environment values."""
|
|
|
|
values = read_env_file(Path(env_path))
|
|
values.update(os.environ if environ is None else environ)
|
|
|
|
base_url = values.get(ORACLE_BASE_URL_VARIABLE, "").strip() or None
|
|
model = values.get(EMBEDDING_MODEL_VARIABLE, "").strip() or None
|
|
dimension = parse_optional_dimension(values.get(EMBEDDING_DIMENSION_VARIABLE))
|
|
return EmbeddingConfig(base_url, model, dimension)
|
|
|
|
|
|
def parse_optional_dimension(value: str | None) -> int | None:
|
|
"""Interpret a blank dimension as unknown and validate configured values."""
|
|
|
|
if value is None or not value.strip():
|
|
return None
|
|
try:
|
|
dimension = int(value)
|
|
except ValueError as error:
|
|
raise ValueError(
|
|
f"{EMBEDDING_DIMENSION_VARIABLE} must be a positive integer or blank"
|
|
) from error
|
|
if dimension <= 0:
|
|
raise ValueError(
|
|
f"{EMBEDDING_DIMENSION_VARIABLE} must be a positive integer or blank"
|
|
)
|
|
return dimension
|