core/config: add database section to config store

Parse and expose [database] with defaults; keep providers API unchanged.
This commit is contained in:
2026-07-14 17:19:21 +08:00
parent 7f63e7980b
commit c3e44cfff0
6 changed files with 77 additions and 12 deletions
+1
View File
@@ -0,0 +1 @@
from . import config as config
+1
View File
@@ -6,6 +6,7 @@ import tomlkit
from tomlkit.exceptions import TOMLKitError from tomlkit.exceptions import TOMLKitError
from .models import AnthropicProvider as AnthropicProvider from .models import AnthropicProvider as AnthropicProvider
from .models import DatabaseConfig as DatabaseConfig
from .models import DeepseekProvider as DeepseekProvider from .models import DeepseekProvider as DeepseekProvider
from .models import ModelConfig as ModelConfig from .models import ModelConfig as ModelConfig
from .models import OpenAICompatibleProvider as OpenAICompatibleProvider from .models import OpenAICompatibleProvider as OpenAICompatibleProvider
+9
View File
@@ -1,6 +1,15 @@
from msgspec import Struct, field from msgspec import Struct, field
class DatabaseConfig(Struct, omit_defaults=True):
"""Database connection configuration."""
implementation: str = "sqlite"
url: str = ":memory:"
username: str | None = None
password: str | None = None
class ModelConfig(Struct, omit_defaults=True): class ModelConfig(Struct, omit_defaults=True):
"""Capability flags for a model within a provider.""" """Capability flags for a model within a provider."""
+52 -12
View File
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, cast
import msgspec import msgspec
import tomlkit import tomlkit
from .models import Provider from .models import DatabaseConfig, Provider
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Mapping from collections.abc import Mapping
@@ -17,6 +17,14 @@ class ConfigFormatError(ValueError):
"""Raised when the config file contains invalid TOML.""" """Raised when the config file contains invalid TOML."""
def _parse_database(raw: dict[str, object]) -> DatabaseConfig:
"""Parse the ``[database]`` section, falling back to defaults."""
try:
return msgspec.convert(raw, DatabaseConfig)
except msgspec.ValidationError:
return DatabaseConfig()
def _parse_providers( def _parse_providers(
document: tomlkit.TOMLDocument, document: tomlkit.TOMLDocument,
) -> tuple[dict[str, Provider], dict[str, object]]: ) -> tuple[dict[str, Provider], dict[str, object]]:
@@ -64,14 +72,26 @@ class ConfigStore:
_path: Path _path: Path
_document: tomlkit.TOMLDocument _document: tomlkit.TOMLDocument
_database: DatabaseConfig
_providers: dict[str, Provider] _providers: dict[str, Provider]
_bad_providers: dict[str, object] _bad_providers: dict[str, object]
def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None: def __init__(self, path: Path, document: tomlkit.TOMLDocument) -> None:
self._path = path self._path = path
self._document = document self._document = document
raw: dict[str, object] = document.unwrap()
self._database = _parse_database(
cast("dict[str, object]", raw.get("database", {}))
)
self._providers, self._bad_providers = _parse_providers(document) self._providers, self._bad_providers = _parse_providers(document)
# -- database (read-only) --
@property
def database(self) -> MappingProxyType[str, object]:
"""Read-only mapping view of database configuration."""
return MappingProxyType(msgspec.structs.asdict(self._database))
# -- providers (read/write) -- # -- providers (read/write) --
@property @property
@@ -98,43 +118,63 @@ class ConfigStore:
# -- persistence -- # -- persistence --
def write(self) -> None: def write(self) -> None:
"""Serialize current providers to the TOML file.""" """Serialize current state to the TOML file."""
self._sync_to_document() self._sync_to_document()
self._path.parent.mkdir(parents=True, exist_ok=True) self._path.parent.mkdir(parents=True, exist_ok=True)
with self._path.open("w") as f: with self._path.open("w") as f:
tomlkit.dump(self._document, f) tomlkit.dump(self._document, f)
def reload(self) -> None: def reload(self) -> None:
"""Re-read the TOML file and re-parse providers.""" """Re-read the TOML file and re-parse all sections."""
with self._path.open() as f: with self._path.open() as f:
self._document = tomlkit.parse(f.read()) self._document = tomlkit.parse(f.read())
raw: dict[str, object] = self._document.unwrap()
self._database = _parse_database(
cast("dict[str, object]", raw.get("database", {}))
)
self._providers, self._bad_providers = _parse_providers(self._document) self._providers, self._bad_providers = _parse_providers(self._document)
# -- internal -- # -- internal sync helpers --
def _sync_to_document(self) -> None: def _sync_database_section(self) -> None:
"""Incrementally sync ``self._providers`` into the ``[providers]`` section. """Sync ``[database]`` to the document."""
raw_db: dict[str, object] = msgspec.to_builtins(self._database)
if not raw_db:
if "database" in self._document:
del self._document["database"]
return
Only adds, updates, or removes individual provider entries — the rest of if "database" not in self._document:
the document (comments, formatting, other top-level sections) is untouched. self._document["database"] = tomlkit.table()
""" section = self._document["database"]
for k in list(section.keys()): # type: ignore[union-attr]
if k not in raw_db:
del section[k] # type: ignore[union-attr]
for k, v in raw_db.items():
section[k] = v # type: ignore[union-attr]
def _sync_providers_section(self) -> None:
"""Sync ``[providers]`` to the document."""
if "providers" not in self._document: if "providers" not in self._document:
self._document["providers"] = tomlkit.table() self._document["providers"] = tomlkit.table()
section = self._document["providers"] section = self._document["providers"]
# Remove deleted providers
for name in list(section.keys()): # type: ignore[union-attr] for name in list(section.keys()): # type: ignore[union-attr]
if name not in self._providers: if name not in self._providers:
del section[name] # type: ignore[union-attr] del section[name] # type: ignore[union-attr]
# Add / update providers
for name, provider in self._providers.items(): for name, provider in self._providers.items():
raw: dict[str, object] = msgspec.to_builtins(provider) raw: dict[str, object] = msgspec.to_builtins(provider)
if name in section: if name in section:
entry = section[name] # type: ignore[union-attr] entry = section[name] # type: ignore[union-attr]
# Clear and rebuild the existing entry to match current state
entry.clear() # type: ignore[union-attr] entry.clear() # type: ignore[union-attr]
for k, v in raw.items(): for k, v in raw.items():
entry[k] = v # type: ignore[union-attr] entry[k] = v # type: ignore[union-attr]
else: else:
section[name] = raw # type: ignore[union-attr] section[name] = raw # type: ignore[union-attr]
def _sync_to_document(self) -> None:
"""Incrementally sync all sections into the document."""
self._sync_database_section()
self._sync_providers_section()
+4
View File
@@ -1,3 +1,7 @@
[database]
implementation = "sqlite"
url = ":memory:"
[providers.test1] [providers.test1]
preset = "openai" preset = "openai"
access_key_or_token = "sk-1145141919810" access_key_or_token = "sk-1145141919810"
+10
View File
@@ -28,6 +28,11 @@ def test_read_default_config(default_config_source: None) -> None:
assert isinstance(providers["test2"], OpenAICompatibleProvider) assert isinstance(providers["test2"], OpenAICompatibleProvider)
assert isinstance(providers["test3"], AnthropicProvider) assert isinstance(providers["test3"], AnthropicProvider)
assert isinstance(providers["foo1"], DeepseekProvider) assert isinstance(providers["foo1"], DeepseekProvider)
db = config.database
assert db["implementation"] == "sqlite"
assert db["url"] == ":memory:"
assert db["username"] is None
assert db["password"] is None
def test_read_valid_config() -> None: def test_read_valid_config() -> None:
@@ -38,6 +43,11 @@ def test_read_valid_config() -> None:
assert isinstance(providers["test2"], OpenAICompatibleProvider) assert isinstance(providers["test2"], OpenAICompatibleProvider)
assert isinstance(providers["test3"], AnthropicProvider) assert isinstance(providers["test3"], AnthropicProvider)
assert isinstance(providers["foo1"], DeepseekProvider) assert isinstance(providers["foo1"], DeepseekProvider)
db = config.database
assert db["implementation"] == "sqlite"
assert db["url"] == ":memory:"
assert db["username"] is None
assert db["password"] is None
def test_read_empty_config() -> None: def test_read_empty_config() -> None: