mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/config: add database section to config store
Parse and expose [database] with defaults; keep providers API unchanged.
This commit is contained in:
@@ -0,0 +1 @@
|
|||||||
|
from . import config as config
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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."""
|
||||||
|
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user