diff --git a/migrations/versions/a1b2c3d4e5f6_add_settings_table.py b/migrations/versions/a1b2c3d4e5f6_add_settings_table.py new file mode 100644 index 00000000..f3c898b8 --- /dev/null +++ b/migrations/versions/a1b2c3d4e5f6_add_settings_table.py @@ -0,0 +1,35 @@ +"""add settings table + +Revision ID: a1b2c3d4e5f6 +Revises: 042f6b77d69d +Create Date: 2025-09-06 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "a1b2c3d4e5f6" +down_revision = "042f6b77d69d" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "settings", + sa.Column("id", sa.Integer(), primary_key=True, nullable=False), + sa.Column("data", sa.Text(), nullable=False), + sa.Column( + "updated_at", + sa.DateTime(), + nullable=True, + server_default=sa.text("CURRENT_TIMESTAMP"), + ), + ) + + +def downgrade() -> None: + op.drop_table("settings") diff --git a/routstr/core/settings.py b/routstr/core/settings.py new file mode 100644 index 00000000..9e8c1d0e --- /dev/null +++ b/routstr/core/settings.py @@ -0,0 +1,293 @@ +from __future__ import annotations + +import asyncio +import json +import os +from datetime import datetime, timezone +from typing import Any + +from pydantic.v1 import BaseModel, BaseSettings, Field +from sqlmodel.ext.asyncio.session import AsyncSession + + +class Settings(BaseSettings): + class Config: + case_sensitive = True + + @classmethod + def parse_env_var(cls, field_name: str, raw_value: str) -> Any: # type: ignore[override] + if field_name in {"cashu_mints", "cors_origins", "relays"}: + v = str(raw_value).strip() + if v == "": + return [] + return [p.strip() for p in v.split(",") if p.strip()] + return raw_value + + # Core + upstream_base_url: str = Field(default="", env="UPSTREAM_BASE_URL") + upstream_api_key: str = Field(default="", env="UPSTREAM_API_KEY") + admin_password: str = Field(default="", env="ADMIN_PASSWORD") + + # Node info + name: str = Field(default="ARoutstrNode", env="NAME") + description: str = Field(default="A Routstr Node", env="DESCRIPTION") + npub: str = Field(default="", env="NPUB") + http_url: str = Field(default="", env="HTTP_URL") + onion_url: str = Field(default="", env="ONION_URL") + + # Cashu + cashu_mints: list[str] = Field(default_factory=list, env="CASHU_MINTS") + receive_ln_address: str = Field(default="", env="RECEIVE_LN_ADDRESS") + primary_mint: str = Field(default="", env="PRIMARY_MINT_URL") + + # Pricing + # Default behavior: derive pricing from MODELS + # If fixed_pricing is True -> use fixed_cost_per_request and ignore tokens + # If fixed_per_1k_* are set (non-zero) -> override model token pricing when model-based + fixed_pricing: bool = Field(default=False, env="FIXED_PRICING") + fixed_cost_per_request: int = Field(default=1, env="FIXED_COST_PER_REQUEST") + fixed_per_1k_input_tokens: int = Field(default=0, env="FIXED_PER_1K_INPUT_TOKENS") + fixed_per_1k_output_tokens: int = Field(default=0, env="FIXED_PER_1K_OUTPUT_TOKENS") + exchange_fee: float = Field(default=1.005, env="EXCHANGE_FEE") + upstream_provider_fee: float = Field(default=1.05, env="UPSTREAM_PROVIDER_FEE") + + # Network + cors_origins: list[str] = Field(default_factory=lambda: ["*"], env="CORS_ORIGINS") + tor_proxy_url: str = Field(default="socks5://127.0.0.1:9050", env="TOR_PROXY_URL") + providers_refresh_interval_seconds: int = Field( + default=300, env="PROVIDERS_REFRESH_INTERVAL_SECONDS" + ) + refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS") + + # Logging + log_level: str = Field(default="INFO", env="LOG_LEVEL") + enable_console_logging: bool = Field(default=True, env="ENABLE_CONSOLE_LOGGING") + + # Other + chat_completions_api_version: str = Field( + default="", env="CHAT_COMPLETIONS_API_VERSION" + ) + models_path: str = Field(default="models.json", env="MODELS_PATH") + source: str = Field(default="", env="SOURCE") + openrouter_base_url: str = Field( + default="https://openrouter.ai/api/v1", env="BASE_URL" + ) + + # Secrets / optional runtime controls + provider_id: str = Field(default="", env="PROVIDER_ID") + nip91_provider_id: str = Field(default="", env="NIP91_PROVIDER_ID") + nsec: str = Field(default="", env="NSEC") + + # NIP-91 + relays: list[str] = Field(default_factory=list, env="RELAYS") + nip91_backoff_base_seconds: float = Field( + default=5.0, env="NIP91_BACKOFF_BASE_SECONDS" + ) + nip91_backoff_max_seconds: float = Field( + default=900.0, env="NIP91_BACKOFF_MAX_SECONDS" + ) + nip91_backoff_jitter_ratio: float = Field( + default=0.2, env="NIP91_BACKOFF_JITTER_RATIO" + ) + nip91_announcement_interval: int = Field( + default=24 * 60 * 60, env="NIP91_ANNOUNCEMENT_INTERVAL" + ) + + +def _compute_primary_mint(cashu_mints: list[str]) -> str: + return cashu_mints[0] if cashu_mints else "https://mint.minibits.cash/Bitcoin" + + +def resolve_bootstrap() -> Settings: + base = Settings() # Reads env with custom parse_env_var + # Back-compat env mapping + try: + # Map MODEL_BASED_PRICING -> fixed_pricing (inverted) + if "MODEL_BASED_PRICING" in os.environ and "FIXED_PRICING" not in os.environ: + mbp_raw = os.environ.get("MODEL_BASED_PRICING", "").strip().lower() + mbp = mbp_raw in {"1", "true", "yes", "on"} + base.fixed_pricing = not mbp + # Map COST_PER_REQUEST -> fixed_cost_per_request if new not provided + if ( + "COST_PER_REQUEST" in os.environ + and "FIXED_COST_PER_REQUEST" not in os.environ + ): + try: + base.fixed_cost_per_request = int( + os.environ["COST_PER_REQUEST"].strip() + ) + except Exception: + pass + # Map COST_PER_1K_* -> CUSTOM_PER_1K_* + if ( + "COST_PER_1K_INPUT_TOKENS" in os.environ + and "FIXED_PER_1K_INPUT_TOKENS" not in os.environ + ): + try: + base.fixed_per_1k_input_tokens = int( + os.environ["COST_PER_1K_INPUT_TOKENS"].strip() + ) + except Exception: + pass + if ( + "COST_PER_1K_OUTPUT_TOKENS" in os.environ + and "FIXED_PER_1K_OUTPUT_TOKENS" not in os.environ + ): + try: + base.fixed_per_1k_output_tokens = int( + os.environ["COST_PER_1K_OUTPUT_TOKENS"].strip() + ) + except Exception: + pass + except Exception: + pass + if not base.onion_url: + try: + from ..nip91 import discover_onion_url_from_tor # type: ignore + + discovered = discover_onion_url_from_tor() + if discovered: + base.onion_url = discovered + except Exception: + pass + if not base.cors_origins: + base.cors_origins = ["*"] + if not base.primary_mint: + base.primary_mint = _compute_primary_mint(base.cashu_mints) + return base + + +class SettingsRow(BaseModel): + id: int + data: dict[str, Any] + updated_at: datetime | None = None + + +# Single, concrete settings instance that callers import directly +settings: Settings = resolve_bootstrap() + + +class SettingsService: + _current: Settings | None = None + _lock: asyncio.Lock = asyncio.Lock() + + @classmethod + def get(cls) -> Settings: + if cls._current is None: + raise RuntimeError("SettingsService not initialized") + return cls._current + + @classmethod + async def initialize(cls, db_session: AsyncSession) -> Settings: + async with cls._lock: + from sqlmodel import text + + await db_session.exec( # type: ignore + text( + "CREATE TABLE IF NOT EXISTS settings (id INTEGER PRIMARY KEY, data TEXT NOT NULL, updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP)" + ) + ) + + row = await db_session.exec( # type: ignore + text("SELECT id, data, updated_at FROM settings WHERE id = 1") + ) + row = row.first() + env_resolved = resolve_bootstrap() + + if row is None: + await db_session.exec( # type: ignore + text( + "INSERT INTO settings (id, data, updated_at) VALUES (1, :data, :updated_at)" + ).bindparams( + data=json.dumps(env_resolved.dict()), + updated_at=datetime.now(timezone.utc), + ) + ) + await db_session.commit() + cls._current = settings + # Update the existing instance in-place for all live importers + for k, v in env_resolved.dict().items(): + setattr(settings, k, v) + return cls._current + + db_id, db_data, _updated_at = row + try: + db_json = ( + json.loads(db_data) if isinstance(db_data, str) else dict(db_data) + ) + except Exception: + db_json = {} + + merged_dict: dict[str, Any] = dict(env_resolved.dict()) + merged_dict.update( + {k: v for k, v in db_json.items() if v not in (None, "")} + ) + + # Ensure primary_mint is consistent with cashu_mints if not explicitly set + if not merged_dict.get("primary_mint"): + merged_dict["primary_mint"] = _compute_primary_mint( + merged_dict.get("cashu_mints", []) + ) + + if any(k not in db_json for k in merged_dict.keys()): + await db_session.exec( # type: ignore + text( + "UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1" + ).bindparams( + data=json.dumps(merged_dict), + updated_at=datetime.now(timezone.utc), + ) + ) + await db_session.commit() + + # Update the existing instance in-place for all live importers + for k, v in merged_dict.items(): + setattr(settings, k, v) + cls._current = settings + return cls._current + + @classmethod + async def update( + cls, partial: dict[str, Any], db_session: AsyncSession + ) -> Settings: + async with cls._lock: + current = cls.get() + candidate_dict = {**current.dict(), **partial} + candidate = Settings(**candidate_dict) + from sqlmodel import text + + # Ensure primary_mint reflects candidate mints if missing + if not candidate.primary_mint: + candidate.primary_mint = _compute_primary_mint(candidate.cashu_mints) + + await db_session.exec( # type: ignore + text( + "UPDATE settings SET data = :data, updated_at = :updated_at WHERE id = 1" + ).bindparams( + data=json.dumps(candidate.dict()), + updated_at=datetime.now(timezone.utc), + ) + ) + await db_session.commit() + # Update in-place + for k, v in candidate.dict().items(): + setattr(settings, k, v) + cls._current = settings + return settings + + @classmethod + async def reload_from_db(cls, db_session: AsyncSession) -> Settings: + async with cls._lock: + from sqlmodel import text + + row = await db_session.exec(text("SELECT data FROM settings WHERE id = 1")) # type: ignore + row = row.first() + if row is None: + raise RuntimeError("Settings row missing") + (data_str,) = row + data = json.loads(data_str) if isinstance(data_str, str) else dict(data_str) + # Update in-place + for k, v in data.items(): + setattr(settings, k, v) + cls._current = settings + return settings diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py new file mode 100644 index 00000000..5c5cb048 --- /dev/null +++ b/tests/unit/test_settings.py @@ -0,0 +1,37 @@ +import os + +import pytest +from sqlalchemy.ext.asyncio import create_async_engine +from sqlmodel.ext.asyncio.session import AsyncSession + +from routstr.core.settings import SettingsService + + +@pytest.mark.asyncio +async def test_settings_seed_from_env_and_persist() -> None: + os.environ["UPSTREAM_BASE_URL"] = "https://api.test/v1" + os.environ.pop("ONION_URL", None) + + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with AsyncSession(engine, expire_on_commit=False) as session: + settings = await SettingsService.initialize(session) + + assert settings.upstream_base_url == "https://api.test/v1" + # ONION_URL may be empty if not discoverable + assert isinstance(settings.onion_url, str) + + +@pytest.mark.asyncio +async def test_settings_db_precedence_over_env() -> None: + os.environ["UPSTREAM_BASE_URL"] = "https://api.env/v1" + + engine = create_async_engine("sqlite+aiosqlite:///:memory:") + async with AsyncSession(engine, expire_on_commit=False) as session: + _ = await SettingsService.initialize(session) + updated = await SettingsService.update({"name": "DBName"}, session) + assert updated.name == "DBName" + + # Change env and re-initialize; DB should still win + os.environ["NAME"] = "EnvName" + again = await SettingsService.initialize(session) + assert again.name == "DBName"