From b0b8ea95f221791ccfcf6bf98c7b121f04508707 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?dni=20=E2=9A=A1?= Date: Wed, 25 Mar 2026 10:47:13 +0100 Subject: [PATCH] 2nd stage draft pydantic --- lnbits/commands.py | 4 +- lnbits/core/crud/settings.py | 2 +- lnbits/core/models/assets.py | 2 +- lnbits/core/models/audit.py | 6 +- lnbits/core/models/extensions.py | 14 +- lnbits/core/models/extensions_builder.py | 52 +++++--- lnbits/core/models/misc.py | 2 +- lnbits/core/models/notifications.py | 2 +- lnbits/core/models/payments.py | 23 ++-- lnbits/core/models/tinyurl.py | 2 +- lnbits/core/models/users.py | 27 ++-- lnbits/core/models/wallets.py | 18 +-- lnbits/core/models/webpush.py | 2 +- lnbits/core/services/extensions_builder.py | 6 +- lnbits/core/services/settings.py | 3 +- lnbits/core/services/users.py | 6 +- lnbits/core/views/admin_api.py | 4 +- lnbits/core/views/auth_api.py | 6 +- lnbits/core/views/extension_api.py | 3 +- lnbits/core/views/node_api.py | 2 +- lnbits/db.py | 122 +++++++++++++---- lnbits/fiat/base.py | 2 +- lnbits/fiat/paypal.py | 13 +- lnbits/fiat/stripe.py | 21 ++- lnbits/helpers.py | 2 +- lnbits/nodes/base.py | 2 +- lnbits/nodes/lndrest.py | 2 +- lnbits/settings.py | 146 ++++++++++++++++----- lnbits/wallets/blink.py | 2 +- pyproject.toml | 1 + tests/unit/test_core_models.py | 80 +++++++++++ tests/unit/test_settings.py | 60 ++++++++- uv.lock | 16 +++ 33 files changed, 499 insertions(+), 156 deletions(-) create mode 100644 tests/unit/test_core_models.py diff --git a/lnbits/commands.py b/lnbits/commands.py index 87c8ddff9..62e4fec5f 100644 --- a/lnbits/commands.py +++ b/lnbits/commands.py @@ -711,7 +711,9 @@ async def _call_install_extension( user_id = user_id or get_super_user() async with httpx.AsyncClient() as client: resp = await client.post( - f"{url}/api/v1/extension?usr={user_id}", json=data.dict(), timeout=40 + f"{url}/api/v1/extension?usr={user_id}", + json=data.model_dump(), + timeout=40, ) resp.raise_for_status() else: diff --git a/lnbits/core/crud/settings.py b/lnbits/core/crud/settings.py index 9e9023e6f..6ee3f153e 100644 --- a/lnbits/core/crud/settings.py +++ b/lnbits/core/crud/settings.py @@ -44,7 +44,7 @@ async def update_admin_settings( data: EditableSettings, tag: str | None = "core" ) -> None: editable_settings = await get_settings_by_tag("core") or {} - editable_settings.update(data.dict(exclude_unset=True)) + editable_settings.update(data.model_dump(exclude_unset=True)) for key, value in editable_settings.items(): try: await set_settings_field(key, value, tag) diff --git a/lnbits/core/models/assets.py b/lnbits/core/models/assets.py index 05ee447df..59687243d 100644 --- a/lnbits/core/models/assets.py +++ b/lnbits/core/models/assets.py @@ -2,7 +2,7 @@ from __future__ import annotations from datetime import datetime, timezone -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field from lnbits.db import FilterModel diff --git a/lnbits/core/models/audit.py b/lnbits/core/models/audit.py index 2c1ce7c22..101dcc665 100644 --- a/lnbits/core/models/audit.py +++ b/lnbits/core/models/audit.py @@ -1,8 +1,9 @@ from __future__ import annotations from datetime import datetime, timedelta, timezone +from typing import Any -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field from lnbits.db import FilterModel from lnbits.settings import settings @@ -21,8 +22,7 @@ class AuditEntry(BaseModel): delete_at: datetime | None = None created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: retention_days = max(0, settings.lnbits_audit_retention_days) or 365 self.delete_at = self.created_at + timedelta(days=retention_days) diff --git a/lnbits/core/models/extensions.py b/lnbits/core/models/extensions.py index bbe2c439c..59b6f005a 100644 --- a/lnbits/core/models/extensions.py +++ b/lnbits/core/models/extensions.py @@ -12,7 +12,7 @@ from typing import Any import httpx from loguru import logger -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field from lnbits.helpers import ( download_url, @@ -96,7 +96,7 @@ class ExtensionConfig(BaseModel): ) error_msg = "Cannot fetch GitHub extension config" config = await github_api_get(config_url, error_msg) - return ExtensionConfig.parse_obj(config) + return ExtensionConfig.model_validate(config) class ReleasePaymentInfo(BaseModel): @@ -304,7 +304,7 @@ class ExtensionRelease(BaseModel): releases_url = f"https://api.github.com/repos/{org}/{repo}/releases" error_msg = "Cannot fetch extension releases" releases = await github_api_get(releases_url, error_msg) - return [GitHubRepoRelease.parse_obj(r) for r in releases] + return [GitHubRepoRelease.model_validate(r) for r in releases] @classmethod async def fetch_release_details(cls, details_link: str) -> dict | None: @@ -757,7 +757,7 @@ class InstallableExtension(BaseModel): repo_url = f"https://api.github.com/repos/{org}/{repository}" error_msg = "Cannot fetch extension repo" repo = await github_api_get(repo_url, error_msg) - github_repo = GitHubRepo.parse_obj(repo) + github_repo = GitHubRepo.model_validate(repo) lates_release_url = ( f"https://api.github.com/repos/{org}/{repository}/releases/latest" @@ -771,15 +771,15 @@ class InstallableExtension(BaseModel): return ( github_repo, - GitHubRepoRelease.parse_obj(latest_release), - ExtensionConfig.parse_obj(config), + GitHubRepoRelease.model_validate(latest_release), + ExtensionConfig.model_validate(config), ) @classmethod async def fetch_manifest(cls, url) -> Manifest: error_msg = "Cannot fetch extensions manifest" manifest = await github_api_get(url, error_msg) - return Manifest.parse_obj(manifest) + return Manifest.model_validate(manifest) class CreateExtension(BaseModel): diff --git a/lnbits/core/models/extensions_builder.py b/lnbits/core/models/extensions_builder.py index fabb9386c..3403635ce 100644 --- a/lnbits/core/models/extensions_builder.py +++ b/lnbits/core/models/extensions_builder.py @@ -6,7 +6,7 @@ import uuid from datetime import datetime, timedelta, timezone from typing import Any, Literal -from pydantic.v1 import BaseModel, validator +from pydantic import BaseModel, field_validator from lnbits.helpers import ( camel_to_snake, @@ -141,7 +141,8 @@ class DataField(BaseModel): else: return f"{self.name} {index}" - @validator("name") + @field_validator("name") + @classmethod def validate_name(cls, v: str) -> str: if v.strip() == "": raise ValueError("Field name is required.") @@ -149,7 +150,8 @@ class DataField(BaseModel): raise ValueError(f"Field Name must be snake_case. Found: {v}") return v - @validator("type") + @field_validator("type") + @classmethod def validate_type(cls, v: str) -> str: if v.strip() == "": raise ValueError("Owner Data type is required") @@ -171,7 +173,8 @@ class DataField(BaseModel): ) return v - @validator("label") + @field_validator("label") + @classmethod def validate_label(cls, v: str | None) -> str | None: if v and '"' in v: raise ValueError( @@ -179,7 +182,8 @@ class DataField(BaseModel): ) return v - @validator("hint") + @field_validator("hint") + @classmethod def validate_hint(cls, v: str | None) -> str | None: if v and '"' in v: raise ValueError(f'Field hint cannot contain double quotes ("). Value: {v}') @@ -191,8 +195,7 @@ class DataFields(BaseModel): editable: bool = True fields: list[DataField] = [] - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: self.normalize() def normalize(self) -> None: @@ -210,7 +213,8 @@ class DataFields(BaseModel): return field return None - @validator("name") + @field_validator("name") + @classmethod def validate_name(cls, v: str) -> str: if v.strip() == "": raise ValueError("Data fields name is required") @@ -223,12 +227,13 @@ class SettingsFields(DataFields): enabled: bool = False type: str = "user" - @validator("type") + @field_validator("type") + @classmethod def validate_type(cls, v: str) -> str: if v.strip() == "": raise ValueError("Settings type is required") if v not in ["user", "admin"]: - raise ValueError("Field Type must be one of: user, admin." f" Found: {v}") + raise ValueError(f"Field Type must be one of: user, admin. Found: {v}") return v @@ -265,8 +270,7 @@ class PreviewAction(BaseModel): is_client_data_preview: bool = False is_public_page_preview: bool = False - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: if not self.is_preview_mode: self.is_settings_preview = False self.is_owner_data_preview = False @@ -286,8 +290,7 @@ class ExtensionData(BaseModel): public_page: PublicPageFields preview_action: PreviewAction = PreviewAction() - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: self.validate_data() self.normalize() @@ -434,7 +437,8 @@ class ExtensionData(BaseModel): f" Received: {paid_flag_field.type}." ) - @validator("id") + @field_validator("id") + @classmethod def validate_id(cls, v: str) -> str: if v.strip() == "": raise ValueError("Extension ID is required") @@ -442,13 +446,15 @@ class ExtensionData(BaseModel): raise ValueError(f"Extension Id must be snake_case. Found: {v}") return v - @validator("name") + @field_validator("name") + @classmethod def validate_name(cls, v: str) -> str: if v.strip() == "": raise ValueError("Extension name is required") return v - @validator("stub_version") + @field_validator("stub_version") + @classmethod def validate_stub_version(cls, v: str | None) -> str | None: if v and '"' in v: raise ValueError( @@ -456,7 +462,8 @@ class ExtensionData(BaseModel): ) return v - @validator("short_description") + @field_validator("short_description") + @classmethod def validate_short_description(cls, v: str | None) -> str | None: if v and '"' in v: raise ValueError( @@ -464,7 +471,8 @@ class ExtensionData(BaseModel): ) return v - @validator("description") + @field_validator("description") + @classmethod def validate_description(cls, v: str | None) -> str | None: if v and '"' in v: raise ValueError( @@ -472,13 +480,15 @@ class ExtensionData(BaseModel): ) return v - @validator("owner_data") + @field_validator("owner_data") + @classmethod def validate_owner_data(cls, v: DataFields) -> DataFields: if len(v.fields) == 0: raise ValueError("At least one owner data field is required") return v - @validator("client_data") + @field_validator("client_data") + @classmethod def validate_client_data(cls, v: DataFields) -> DataFields: if len(v.fields) == 0: raise ValueError("At least one client data field is required") diff --git a/lnbits/core/models/misc.py b/lnbits/core/models/misc.py index 966e364d1..5aad3fdc9 100644 --- a/lnbits/core/models/misc.py +++ b/lnbits/core/models/misc.py @@ -2,7 +2,7 @@ from __future__ import annotations from collections.abc import Callable -from pydantic.v1 import BaseModel +from pydantic import BaseModel def _do_nothing(*_): diff --git a/lnbits/core/models/notifications.py b/lnbits/core/models/notifications.py index 57f331f8d..aa10ac91d 100644 --- a/lnbits/core/models/notifications.py +++ b/lnbits/core/models/notifications.py @@ -1,6 +1,6 @@ from enum import Enum -from pydantic.v1 import BaseModel +from pydantic import BaseModel from lnbits.core.models.users import UserNotifications diff --git a/lnbits/core/models/payments.py b/lnbits/core/models/payments.py index 7deae3b0b..b03ba119c 100644 --- a/lnbits/core/models/payments.py +++ b/lnbits/core/models/payments.py @@ -2,13 +2,13 @@ from __future__ import annotations from datetime import datetime, timezone from enum import Enum -from typing import Literal +from typing import Any, Literal from fastapi import Query # from lnurl import LnurlWithdrawResponse from loguru import logger -from pydantic.v1 import BaseModel, Field, validator +from pydantic import BaseModel, Field, field_validator from lnbits.db import FilterModel from lnbits.fiat.base import ( @@ -63,7 +63,10 @@ class Payment(BaseModel): amount: int fee: int bolt11: str - payment_request: str | None = Field(default=None, no_database=True) + payment_request: str | None = Field( + default=None, + json_schema_extra={"no_database": True}, + ) fiat_provider: str | None = None status: str = PaymentState.PENDING memo: str | None = None @@ -79,8 +82,7 @@ class Payment(BaseModel): labels: list[str] = [] extra: dict = {} - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: if "fiat_payment_request" in self.extra: self.payment_request = self.extra["fiat_payment_request"] else: @@ -252,13 +254,14 @@ class CreateInvoice(BaseModel): fiat_provider: str | None = None labels: list[str] = [] - @validator("payment_hash") + @field_validator("payment_hash") + @classmethod def check_hex(cls, v): if v: _ = bytes.fromhex(v) return v - @validator("unit") + @field_validator("unit") @classmethod def unit_is_from_allowed_currencies(cls, v): if v != "sat" and v not in allowed_currencies(): @@ -281,7 +284,8 @@ class SettleInvoice(BaseModel): max_length=64, ) - @validator("preimage") + @field_validator("preimage") + @classmethod def check_hex(cls, v): _ = bytes.fromhex(v) return v @@ -295,7 +299,8 @@ class CancelInvoice(BaseModel): max_length=64, ) - @validator("payment_hash") + @field_validator("payment_hash") + @classmethod def check_hex(cls, v): _ = bytes.fromhex(v) return v diff --git a/lnbits/core/models/tinyurl.py b/lnbits/core/models/tinyurl.py index 9090cd869..a9e4cdff5 100644 --- a/lnbits/core/models/tinyurl.py +++ b/lnbits/core/models/tinyurl.py @@ -1,4 +1,4 @@ -from pydantic.v1 import BaseModel +from pydantic import BaseModel class TinyURL(BaseModel): diff --git a/lnbits/core/models/users.py b/lnbits/core/models/users.py index 938632ad3..cd6d496a4 100644 --- a/lnbits/core/models/users.py +++ b/lnbits/core/models/users.py @@ -1,11 +1,12 @@ from __future__ import annotations from datetime import datetime, timezone +from typing import Any from uuid import UUID from bcrypt import checkpw, gensalt, hashpw from fastapi import Query -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field from lnbits.core.models.misc import SimpleItem from lnbits.db import FilterModel @@ -38,11 +39,9 @@ class WalletInviteRequest(BaseModel): class UserLabel(BaseModel): - name: str = Field(regex=r"([A-Za-z0-9 ._-]{1,100}$)") + name: str = Field(pattern=r"([A-Za-z0-9 ._-]{1,100}$)") description: str | None = Field(default=None, max_length=250) - color: str | None = Field( - default=None, regex=r"^#[0-9A-Fa-f]{6}$" - ) # e.g., "#RRGGBB" + color: str | None = Field(default=None, pattern=r"^#[0-9A-Fa-f]{6}$") class UserExtra(BaseModel): @@ -193,12 +192,20 @@ class Account(AccountId): created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) updated_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) - is_super_user: bool = Field(default=False, no_database=True) - is_admin: bool = Field(default=False, no_database=True) - fiat_providers: list[str] = Field(default=[], no_database=True) + is_super_user: bool = Field( + default=False, + json_schema_extra={"no_database": True}, + ) + is_admin: bool = Field( + default=False, + json_schema_extra={"no_database": True}, + ) + fiat_providers: list[str] = Field( + default=[], + json_schema_extra={"no_database": True}, + ) - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: self.is_super_user = settings.is_super_user(self.id) self.is_admin = settings.is_admin_user(self.id) self.fiat_providers = settings.get_fiat_providers_for_user(self.id) diff --git a/lnbits/core/models/wallets.py b/lnbits/core/models/wallets.py index 3adb7f1e4..689454b66 100644 --- a/lnbits/core/models/wallets.py +++ b/lnbits/core/models/wallets.py @@ -3,10 +3,11 @@ from __future__ import annotations from dataclasses import dataclass from datetime import datetime, timezone from enum import Enum +from typing import Any -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field -# from lnbits.core.models.lnurl import StoredPayLinks +from lnbits.core.models.lnurl import StoredPayLinks from lnbits.db import FilterModel from lnbits.settings import settings @@ -126,15 +127,16 @@ class Wallet(BaseWallet): created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) updated_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc)) currency: str | None = None - balance_msat: int = Field(default=0, no_database=True) + balance_msat: int = Field(default=0, json_schema_extra={"no_database": True}) extra: WalletExtra = WalletExtra() - # TODO dont mix v1 and v2 - # stored_paylinks: StoredPayLinks = StoredPayLinks() + stored_paylinks: StoredPayLinks = Field(default_factory=StoredPayLinks) # What permission this wallet has when it's a shared wallet - share_permissions: list[WalletPermission] = Field(default=[], no_database=True) + share_permissions: list[WalletPermission] = Field( + default=[], + json_schema_extra={"no_database": True}, + ) - def __init__(self, **data): - super().__init__(**data) + def model_post_init(self, __context: Any) -> None: self._validate_data() def mirror_shared_wallet( diff --git a/lnbits/core/models/webpush.py b/lnbits/core/models/webpush.py index f2322e514..b7c35ebf7 100644 --- a/lnbits/core/models/webpush.py +++ b/lnbits/core/models/webpush.py @@ -1,6 +1,6 @@ from datetime import datetime -from pydantic.v1 import BaseModel +from pydantic import BaseModel class CreateWebPushSubscription(BaseModel): diff --git a/lnbits/core/services/extensions_builder.py b/lnbits/core/services/extensions_builder.py index 8a323c91a..f1c899481 100644 --- a/lnbits/core/services/extensions_builder.py +++ b/lnbits/core/services/extensions_builder.py @@ -97,7 +97,7 @@ def _transform_extension_builder_stub(data: ExtensionData, extension_dir: Path) def _export_extension_data_json(data: ExtensionData, build_dir: Path): json.dump( - data.dict(), + data.model_dump(), open(Path(build_dir, "builder.json"), "w", encoding="utf-8"), indent=4, ) @@ -133,7 +133,7 @@ async def _get_extension_stub_release( logger.debug(f"Save release cache {stub_ext_id} ({stub_version}).") with open(release_cache_file, "w", encoding="utf-8") as f: - f.write(json.dumps(release.dict(), indent=4)) + f.write(json.dumps(release.model_dump(), indent=4)) return release @@ -290,7 +290,7 @@ def _replace_jinja_placeholders(data: ExtensionData, ext_stub_dir: Path) -> None { "extension_builder_stub_public_client_inputs": public_client_inputs, "preview": data.preview_action, - **data.public_page.action_fields.dict(), + **data.public_page.action_fields.model_dump(), "cancel_comment": remove_line_marker, }, ) diff --git a/lnbits/core/services/settings.py b/lnbits/core/services/settings.py index ed7b8e859..bc8230cd3 100644 --- a/lnbits/core/services/settings.py +++ b/lnbits/core/services/settings.py @@ -45,10 +45,11 @@ def dict_to_settings(sets_dict: dict) -> UpdateSettings: def update_cached_settings(sets_dict: dict): editable_settings = dict_to_settings(sets_dict) + settings_keys = settings.model_dump().keys() for key in sets_dict.keys(): if key in readonly_variables: continue - if key not in settings.dict().keys(): + if key not in settings_keys: continue try: value = getattr(editable_settings, key) diff --git a/lnbits/core/services/users.py b/lnbits/core/services/users.py index 37f2c998c..2cd69ac28 100644 --- a/lnbits/core/services/users.py +++ b/lnbits/core/services/users.py @@ -163,7 +163,7 @@ async def check_admin_settings(): # .env super_user overwrites DB super_user settings_db = await update_super_user(settings.super_user) - update_cached_settings(settings_db.dict()) + update_cached_settings(settings_db.model_dump()) # saving superuser to {data_dir}/.super_user file with open(Path(settings.lnbits_data_folder) / ".super_user", "w") as file: @@ -192,8 +192,8 @@ async def init_admin_settings(super_user: str | None = None) -> SuperSettings: await create_account(account) await create_wallet(user_id=account.id) - editable_settings = EditableSettings.from_dict(settings.dict()) - return await create_admin_settings(account.id, editable_settings.dict()) + editable_settings = EditableSettings.from_dict(settings.model_dump()) + return await create_admin_settings(account.id, editable_settings.model_dump()) async def check_register_activation_settings(data: RegisterUser): diff --git a/lnbits/core/views/admin_api.py b/lnbits/core/views/admin_api.py index 6e69d07c8..dad6c4f08 100644 --- a/lnbits/core/views/admin_api.py +++ b/lnbits/core/views/admin_api.py @@ -87,7 +87,7 @@ async def api_update_settings( admin_settings = await get_admin_settings(account.is_super_user) if not admin_settings: raise ValueError("Updated admin settings not found.") - update_cached_settings(admin_settings.dict()) + update_cached_settings(admin_settings.model_dump()) core_app_extra.register_new_ratelimiter() return {"status": "Success"} @@ -99,7 +99,7 @@ async def api_update_settings( async def api_update_settings_partial( data: dict, account: Account = Depends(check_admin) ): - updatable_settings = dict_to_settings({**settings.dict(), **data}) + updatable_settings = dict_to_settings({**settings.model_dump(), **data}) return await api_update_settings(updatable_settings, account) diff --git a/lnbits/core/views/auth_api.py b/lnbits/core/views/auth_api.py index 9fd0d8a73..481f138b9 100644 --- a/lnbits/core/views/auth_api.py +++ b/lnbits/core/views/auth_api.py @@ -587,7 +587,7 @@ def _auth_success_response( payload = AccessTokenPayload( sub=username or "", usr=user_id, email=email, auth_time=int(time()) ) - access_token = create_access_token(data=payload.dict()) + access_token = create_access_token(data=payload.model_dump()) max_age = settings.auth_token_expire_minutes * 60 response = JSONResponse({"access_token": access_token, "token_type": "bearer"}) response.set_cookie( @@ -617,7 +617,7 @@ def _auth_api_token_response( sub=username, api_token_id=api_token_id, auth_time=int(time()) ) return create_access_token( - data=payload.dict(), token_expire_minutes=token_expire_minutes + data=payload.model_dump(), token_expire_minutes=token_expire_minutes ) @@ -625,7 +625,7 @@ def _auth_redirect_response(path: str, user_id: str, email: str) -> RedirectResp payload = AccessTokenPayload( usr=user_id, sub="", email=email, auth_time=int(time()) ) - access_token = create_access_token(data=payload.dict()) + access_token = create_access_token(data=payload.model_dump()) max_age = settings.auth_token_expire_minutes * 60 response = RedirectResponse(path) response.set_cookie( diff --git a/lnbits/core/views/extension_api.py b/lnbits/core/views/extension_api.py index 3192b832c..c26cae06d 100644 --- a/lnbits/core/views/extension_api.py +++ b/lnbits/core/views/extension_api.py @@ -637,7 +637,8 @@ async def create_extension_review( ) -> ExtensionReviewPaymentRequest: async with httpx.AsyncClient() as client: resp = await client.post( - settings.lnbits_extensions_reviews_url + "/reviews", json=data.dict() + settings.lnbits_extensions_reviews_url + "/reviews", + json=data.model_dump(), ) resp.raise_for_status() payment_request = resp.json() diff --git a/lnbits/core/views/node_api.py b/lnbits/core/views/node_api.py index 90d71aadf..d299365ce 100644 --- a/lnbits/core/views/node_api.py +++ b/lnbits/core/views/node_api.py @@ -2,7 +2,7 @@ from http import HTTPStatus import httpx from fastapi import APIRouter, Body, Depends, HTTPException -from pydantic.v1 import BaseModel +from pydantic import BaseModel from lnbits.decorators import check_admin, check_super_user, parse_filters from lnbits.settings import settings diff --git a/lnbits/db.py b/lnbits/db.py index cbf5498b1..4bd136c05 100644 --- a/lnbits/db.py +++ b/lnbits/db.py @@ -8,9 +8,11 @@ import time from contextlib import asynccontextmanager from datetime import datetime, timezone from enum import Enum -from typing import Any, Generic, Literal, TypeVar, get_origin +from types import UnionType +from typing import Any, Generic, Literal, TypeVar, Union, get_args, get_origin from loguru import logger +from pydantic import BaseModel as BaseModelV2 from pydantic.v1 import BaseModel, ValidationError, root_validator from sqlalchemy.ext.asyncio import AsyncConnection, AsyncEngine, create_async_engine from sqlalchemy.sql import text @@ -20,6 +22,7 @@ from lnbits.settings import settings POSTGRES = "POSTGRES" COCKROACH = "COCKROACH" SQLITE = "SQLITE" +PYDANTIC_MODEL_TYPES = (BaseModel, BaseModelV2) DateTrunc = Literal["hour", "day", "month"] sqlite_formats = { @@ -28,6 +31,80 @@ sqlite_formats = { "month": "%Y-%m-01 00:00:00", } + +def _is_pydantic_model_class(model: Any) -> bool: + return isinstance(model, type) and any( + issubclass(model, model_type) for model_type in PYDANTIC_MODEL_TYPES + ) + + +def _is_v2_pydantic_model_class(model: Any) -> bool: + return isinstance(model, type) and issubclass(model, BaseModelV2) + + +def _issubclass(candidate: Any, parent: Any) -> bool: + return isinstance(candidate, type) and issubclass(candidate, parent) + + +def _strip_optional(annotation: Any) -> Any: + origin = get_origin(annotation) + if origin not in {Union, UnionType}: + return annotation + + args = [arg for arg in get_args(annotation) if arg is not type(None)] + if len(args) == 1: + return _strip_optional(args[0]) + return annotation + + +def _model_fields(model: type[Any]) -> dict[str, Any]: + fields = getattr(model, "model_fields", None) + if fields is not None: + return fields + return getattr(model, "__fields__", {}) + + +def _field_annotation(field: Any) -> Any: + annotation = getattr(field, "annotation", None) + if annotation is None: + annotation = getattr(field, "outer_type_", Any) + return _strip_optional(annotation) + + +def _field_inner_type(field: Any) -> Any: + type_ = getattr(field, "type_", None) + if type_ is not None: + return _strip_optional(type_) + + annotation = _field_annotation(field) + origin = get_origin(annotation) + if origin in {list, set, tuple, dict}: + args = get_args(annotation) + return _strip_optional(args[0]) if args else Any + + return annotation + + +def _field_extra(field: Any) -> dict[str, Any]: + field_info = getattr(field, "field_info", None) + if field_info is not None: + return field_info.extra + + return getattr(field, "json_schema_extra", None) or {} + + +def _model_dump(model: BaseModel | BaseModelV2) -> dict[str, Any]: + if isinstance(model, BaseModelV2): + return model.model_dump() + return model.dict() + + +def _validate_model(model: type[Any], values: dict[str, Any]) -> Any: + if _is_v2_pydantic_model_class(model): + return model.model_validate(values) + return model.parse_obj(values) + + if settings.lnbits_database_url: database_uri = settings.lnbits_database_url if database_uri.startswith("cockroachdb://"): @@ -35,7 +112,7 @@ if settings.lnbits_database_url: else: if not database_uri.startswith("postgres://"): raise ValueError( - "Please use the 'postgres://...' " "format for the database URL." + "Please use the 'postgres://...' format for the database URL." ) DB_TYPE = POSTGRES @@ -452,7 +529,7 @@ class FilterModel(BaseModel): T = TypeVar("T") -TModel = TypeVar("TModel", bound=BaseModel) +TModel = TypeVar("TModel", bound=BaseModel | BaseModelV2) TFilterModel = TypeVar("TFilterModel", bound=FilterModel) @@ -660,17 +737,19 @@ def update_query( return f"UPDATE {table_name} SET {query} {where}" # noqa: S608 -def model_to_dict(model: BaseModel) -> dict: +def model_to_dict(model: BaseModel | BaseModelV2) -> dict: """ Convert a Pydantic model to a dictionary with JSON-encoded nested models private fields starting with _ are ignored :param model: Pydantic model """ _dict: dict = {} - for key, value in model.dict().items(): - type_ = model.__fields__[key].type_ - outertype_ = model.__fields__[key].outer_type_ - if model.__fields__[key].field_info.extra.get("no_database", False): + model_fields = _model_fields(type(model)) + for key, value in _model_dump(model).items(): + field = model_fields[key] + type_ = _field_inner_type(field) + outertype_ = _field_annotation(field) + if _field_extra(field).get("no_database", False): continue if isinstance(value, datetime): if DB_TYPE == SQLITE: @@ -681,7 +760,7 @@ def model_to_dict(model: BaseModel) -> dict: _dict[key] = value.replace(tzinfo=None) continue if ( - type(type_) is type(BaseModel) + _is_pydantic_model_class(type_) or type_ is dict or get_origin(outertype_) is list ): @@ -705,51 +784,50 @@ def dict_to_submodel(model: type[TModel], value: dict | str) -> TModel | None: return dict_to_model(_subdict, model) -def dict_to_model(_row: dict, model: type[TModel]) -> TModel: # noqa: C901 +def dict_to_model(_row: dict, model: type[TModel]) -> TModel: """ Convert a dictionary with JSON-encoded nested models to a Pydantic model :param _dict: Dictionary from database :param model: Pydantic model """ _dict: dict = {} + model_fields = _model_fields(model) for key, value in _row.items(): if value is None: continue - if key not in model.__fields__: + if key not in model_fields: # Somethimes an SQL JOIN will create additional column continue - type_ = model.__fields__[key].type_ - outertype_ = model.__fields__[key].outer_type_ + field = model_fields[key] + type_ = _field_inner_type(field) + outertype_ = _field_annotation(field) if get_origin(outertype_) is list: _items = _safe_load_json(value) if isinstance(value, str) else value _dict[key] = [ - dict_to_submodel(type_, v) if issubclass(type_, BaseModel) else v + dict_to_submodel(type_, v) if _is_pydantic_model_class(type_) else v for v in _items ] continue - if issubclass(type_, bool): + if _issubclass(type_, bool): _dict[key] = bool(value) continue - if issubclass(type_, datetime): + if _issubclass(type_, datetime): if DB_TYPE == SQLITE: _dict[key] = datetime.fromtimestamp(value, timezone.utc) else: _dict[key] = value.replace(tzinfo=timezone.utc) continue - if issubclass(type_, BaseModel): + if _is_pydantic_model_class(type_): _dict[key] = dict_to_submodel(type_, value) continue # TODO: remove this when all sub models are migrated to Pydantic # NOTE: this is for type dict on BaseModel, (used in Payment class) - if type_ is dict and value: + if (type_ is dict or get_origin(outertype_) is dict) and value: _dict[key] = _safe_load_json(value) continue _dict[key] = value continue - _model = model.construct(**_dict) - if isinstance(_model, BaseModel): - _model.__init__(**_dict) # type: ignore - return _model + return _validate_model(model, _dict) def _safe_load_json(value: str) -> dict: diff --git a/lnbits/fiat/base.py b/lnbits/fiat/base.py index 3b22bb080..162ed85bb 100644 --- a/lnbits/fiat/base.py +++ b/lnbits/fiat/base.py @@ -4,7 +4,7 @@ from abc import ABC, abstractmethod from collections.abc import AsyncGenerator, Coroutine from typing import TYPE_CHECKING, Any, NamedTuple -from pydantic.v1 import BaseModel, Field +from pydantic import BaseModel, Field if TYPE_CHECKING: pass diff --git a/lnbits/fiat/paypal.py b/lnbits/fiat/paypal.py index 39a563df5..d23f5b8f9 100644 --- a/lnbits/fiat/paypal.py +++ b/lnbits/fiat/paypal.py @@ -6,7 +6,7 @@ from typing import Any, Literal import httpx from loguru import logger -from pydantic.v1 import BaseModel, Field, ValidationError +from pydantic import BaseModel, ConfigDict, Field, ValidationError from lnbits.helpers import normalize_endpoint, urlsafe_short_hash from lnbits.settings import settings @@ -28,8 +28,7 @@ FiatMethod = Literal["checkout", "subscription"] class PayPalCheckoutOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") success_url: str | None = None cancel_url: str | None = None @@ -37,16 +36,14 @@ class PayPalCheckoutOptions(BaseModel): class PayPalSubscriptionOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") checking_id: str | None = None payment_request: str | None = None class PayPalCreateInvoiceOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") fiat_method: FiatMethod = "checkout" checkout: PayPalCheckoutOptions | None = None @@ -339,7 +336,7 @@ class PayPalWallet(FiatProvider): self, raw_opts: dict[str, Any] ) -> PayPalCreateInvoiceOptions | None: try: - return PayPalCreateInvoiceOptions.parse_obj(raw_opts) + return PayPalCreateInvoiceOptions.model_validate(raw_opts) except ValidationError as e: logger.warning(f"Invalid PayPal options: {e}") return None diff --git a/lnbits/fiat/stripe.py b/lnbits/fiat/stripe.py index 950f43b51..cff0e172f 100644 --- a/lnbits/fiat/stripe.py +++ b/lnbits/fiat/stripe.py @@ -8,7 +8,7 @@ from urllib.parse import urlencode import httpx from loguru import logger -from pydantic.v1 import BaseModel, Field, ValidationError +from pydantic import BaseModel, ConfigDict, Field, ValidationError from lnbits.helpers import normalize_endpoint, urlsafe_short_hash from lnbits.settings import settings @@ -30,8 +30,7 @@ FiatMethod = Literal["checkout", "terminal", "subscription"] class StripeTerminalOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") capture_method: Literal["automatic", "manual"] = "automatic" metadata: dict[str, str] = Field(default_factory=dict) @@ -39,8 +38,7 @@ class StripeTerminalOptions(BaseModel): class StripeCheckoutOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") success_url: str | None = None metadata: dict[str, str] = Field(default_factory=dict) @@ -48,16 +46,14 @@ class StripeCheckoutOptions(BaseModel): class StripeSubscriptionOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") checking_id: str | None = None payment_request: str | None = None class StripeCreateInvoiceOptions(BaseModel): - class Config: - extra = "ignore" + model_config = ConfigDict(extra="ignore") fiat_method: FiatMethod = "checkout" terminal: StripeTerminalOptions | None = None @@ -171,7 +167,10 @@ class StripeWallet(FiatProvider): ("line_items[0][price]", subscription_id), ("line_items[0][quantity]", f"{quantity}"), ] - subscription_data = {**payment_options.dict(), "alan_action": "subscription"} + subscription_data = { + **payment_options.model_dump(), + "alan_action": "subscription", + } subscription_data["extra"] = json.dumps(subscription_data.get("extra") or {}) form_data += self._encode_metadata( @@ -494,7 +493,7 @@ class StripeWallet(FiatProvider): self, raw_opts: dict[str, Any] ) -> StripeCreateInvoiceOptions | None: try: - return StripeCreateInvoiceOptions.parse_obj(raw_opts) + return StripeCreateInvoiceOptions.model_validate(raw_opts) except ValidationError as e: logger.warning(f"Invalid Stripe options: {e}") return None diff --git a/lnbits/helpers.py b/lnbits/helpers.py index 16ad3dbb7..7d7a9f463 100644 --- a/lnbits/helpers.py +++ b/lnbits/helpers.py @@ -71,7 +71,7 @@ def template_renderer(additional_folders: list | None = None) -> Jinja2Templates # used in base.html t.env.globals["SITE_TITLE"] = settings.lnbits_site_title t.env.globals["LNBITS_APPLE_TOUCH_ICON"] = settings.lnbits_apple_touch_icon - t.env.globals["SETTINGS"] = settings.to_public().dict(by_alias=True) + t.env.globals["SETTINGS"] = settings.to_public().model_dump(by_alias=True) t.env.globals["CURRENCIES"] = list(currencies.keys()) if settings.bundle_assets: diff --git a/lnbits/nodes/base.py b/lnbits/nodes/base.py index 7d12b36aa..a75f822f0 100644 --- a/lnbits/nodes/base.py +++ b/lnbits/nodes/base.py @@ -4,7 +4,7 @@ from abc import ABC, abstractmethod from enum import Enum from typing import TYPE_CHECKING -from pydantic.v1 import BaseModel +from pydantic import BaseModel from lnbits.db import FilterModel, Filters, Page from lnbits.utils.cache import cache diff --git a/lnbits/nodes/lndrest.py b/lnbits/nodes/lndrest.py index aec6315f6..a96b8d2e7 100644 --- a/lnbits/nodes/lndrest.py +++ b/lnbits/nodes/lndrest.py @@ -364,7 +364,7 @@ class LndRestNode(Node): fee_report = await self.get("/v1/fees") balance = await self.get("/v1/balance/channels") return NodeInfoResponse( - **public.dict(), + **public.model_dump(), onchain_balance_sat=onchain["total_balance"], onchain_confirmed_sat=onchain["confirmed_balance"], balance_msat=balance["local_balance"]["msat"], diff --git a/lnbits/settings.py b/lnbits/settings.py index 8e9ea25dc..c38b9c0a7 100644 --- a/lnbits/settings.py +++ b/lnbits/settings.py @@ -2,7 +2,6 @@ from __future__ import annotations import importlib import importlib.metadata -import inspect import json import os import re @@ -11,11 +10,12 @@ from enum import Enum from os import path from pathlib import Path from time import gmtime, strftime, time -from typing import Any +from typing import Any, get_args, get_origin from uuid import uuid4 +from dotenv import dotenv_values from loguru import logger -from pydantic.v1 import BaseModel, BaseSettings, Extra, Field, validator +from pydantic import BaseModel, ConfigDict, Field, field_validator def list_parse_fallback(v: str): @@ -29,6 +29,89 @@ def list_parse_fallback(v: str): return [] +def _remove_env_names(schema: dict[str, Any]) -> None: + for prop in schema.get("properties", {}).values(): + prop.pop("env_names", None) + + +def _iter_validation_aliases(validation_alias: Any) -> list[str]: + if isinstance(validation_alias, str): + return [validation_alias] + + choices = getattr(validation_alias, "choices", None) + if not choices: + return [] + + return [choice for choice in choices if isinstance(choice, str)] + + +def _annotation_has_origin(annotation: Any, origins: tuple[type, ...]) -> bool: + origin = get_origin(annotation) + if origin in origins: + return True + + return any(_annotation_has_origin(arg, origins) for arg in get_args(annotation)) + + +def _parse_settings_value(value: Any, annotation: Any) -> Any: + if not isinstance(value, str): + return value + + stripped_value = value.strip() + + if stripped_value.startswith("[") or stripped_value.startswith("{"): + return json.loads(stripped_value) + + if _annotation_has_origin(annotation, (list, set, tuple)): + return list_parse_fallback(value) + + return value + + +class LNbitsBaseSettings(BaseModel): + def __init__(self, **data): + super().__init__(**{**self._settings_data(), **data}) + + @classmethod + def _settings_data(cls) -> dict[str, Any]: + env_file = cls.model_config.get("env_file") + env_file_encoding = cls.model_config.get("env_file_encoding") + case_sensitive = bool(cls.model_config.get("case_sensitive", False)) + + raw_values: dict[str, Any] = {} + if env_file: + dotenv_items = dotenv_values(env_file, encoding=env_file_encoding) + raw_values.update( + {key: value for key, value in dotenv_items.items() if value is not None} + ) + raw_values.update(os.environ) + + lookup_values = ( + raw_values + if case_sensitive + else {str(key).lower(): value for key, value in raw_values.items()} + ) + + settings_values: dict[str, Any] = {} + for field_name, field in cls.model_fields.items(): + lookup_keys = _iter_validation_aliases(field.validation_alias) + if field.alias and field.alias != field_name: + lookup_keys.append(field.alias) + lookup_keys.append(field_name) + + for key in lookup_keys: + lookup_key = key if case_sensitive else key.lower() + if lookup_key not in lookup_values: + continue + + settings_values[field_name] = _parse_settings_value( + lookup_values[lookup_key], field.annotation + ) + break + + return settings_values + + class LNbitsSettings(BaseModel): @classmethod def validate_list(cls, val): @@ -324,7 +407,7 @@ class AssetSettings(LNbitsSettings): "heif", "heics", "text/plain", - "text/json" "text/xml", + "text/jsontext/xml", "application/json", "application/pdf", ] @@ -652,9 +735,9 @@ class BoltzFundingSource(LNbitsSettings): class StrikeFundingSource(LNbitsSettings): strike_api_endpoint: str | None = Field( - default="https://api.strike.me/v1", env="STRIKE_API_ENDPOINT" + default="https://api.strike.me/v1", validation_alias="STRIKE_API_ENDPOINT" ) - strike_api_key: str | None = Field(default=None, env="STRIKE_API_KEY") + strike_api_key: str | None = Field(default=None, validation_alias="STRIKE_API_KEY") class FiatProviderLimits(BaseModel): @@ -971,12 +1054,12 @@ class EditableSettings( KeycloakAuthSettings, OidcAuthSettings, ): - @validator( + @field_validator( "lnbits_admin_users", "lnbits_allowed_users", "lnbits_theme_options", "lnbits_admin_extensions", - pre=True, + mode="before", ) @classmethod def validate_editable_settings(cls, val): @@ -984,21 +1067,21 @@ class EditableSettings( @classmethod def from_dict(cls, d: dict): - return cls( - **{k: v for k, v in d.items() if k in inspect.signature(cls).parameters} - ) + return cls(**{k: v for k, v in d.items() if k in cls.model_fields}) - # fixes openapi.json validation, remove field env_names - class Config: - @staticmethod - def schema_extra(schema: dict[str, Any]) -> None: - for prop in schema.get("properties", {}).values(): - prop.pop("env_names", None) + # Fixes openapi.json validation by removing the v1-only env_names metadata. + model_config = ConfigDict( + populate_by_name=True, + json_schema_extra=_remove_env_names, + ) class UpdateSettings(EditableSettings): - class Config: - extra = Extra.forbid + model_config = ConfigDict( + populate_by_name=True, + extra="forbid", + json_schema_extra=_remove_env_names, + ) class EnvSettings(LNbitsSettings): @@ -1110,7 +1193,7 @@ class TransientSettings(InstalledExtensionsSettings, ExchangeHistorySettings): @classmethod def readonly_fields(cls): - return [f for f in inspect.signature(cls).parameters if not f.startswith("_")] + return [f for f in cls.model_fields if not f.startswith("_")] class ReadOnlySettings( @@ -1125,9 +1208,9 @@ class ReadOnlySettings( def lnbits_extensions_upgrade_path(self) -> str: return str(Path(self.lnbits_data_folder, "upgrades")) - @validator( + @field_validator( "lnbits_allowed_funding_sources", - pre=True, + mode="before", ) @classmethod def validate_readonly_settings(cls, val): @@ -1135,15 +1218,18 @@ class ReadOnlySettings( @classmethod def readonly_fields(cls): - return [f for f in inspect.signature(cls).parameters if not f.startswith("_")] + return [f for f in cls.model_fields if not f.startswith("_")] -class Settings(EditableSettings, ReadOnlySettings, TransientSettings, BaseSettings): - class Config: - env_file = ".env" - env_file_encoding = "utf-8" - case_sensitive = False - json_loads = list_parse_fallback +class Settings( + EditableSettings, ReadOnlySettings, TransientSettings, LNbitsBaseSettings +): + model_config = ConfigDict( + populate_by_name=True, + env_file=".env", + env_file_encoding="utf-8", + case_sensitive=False, + ) def is_user_allowed(self, user_id: str) -> bool: return ( @@ -1348,7 +1434,7 @@ if not settings.user_agent: # printing environment variable for debugging if not settings.lnbits_admin_ui: logger.debug("Environment Settings:") - for key, value in settings.dict(exclude_none=True).items(): + for key, value in settings.model_dump(exclude_none=True).items(): logger.debug(f"{key}: {value}") diff --git a/lnbits/wallets/blink.py b/lnbits/wallets/blink.py index e9ffd57d8..be736086a 100644 --- a/lnbits/wallets/blink.py +++ b/lnbits/wallets/blink.py @@ -5,7 +5,7 @@ from collections.abc import AsyncGenerator import httpx from loguru import logger -from pydantic.v1 import BaseModel +from pydantic import BaseModel from websockets import Subprotocol, connect from lnbits import bolt11 diff --git a/pyproject.toml b/pyproject.toml index e35b4b953..c5b6edb41 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -51,6 +51,7 @@ dependencies = [ "pillow~=12.1.0", "python-dotenv~=1.2.1", "greenlet~=3.3.0", + "pydantic-settings>=2.13.1", ] [project.scripts] diff --git a/tests/unit/test_core_models.py b/tests/unit/test_core_models.py new file mode 100644 index 000000000..6a539b9bb --- /dev/null +++ b/tests/unit/test_core_models.py @@ -0,0 +1,80 @@ +import pytest +from pydantic import ValidationError + +from lnbits.core.models.extensions import InstallableExtension +from lnbits.core.models.lnurl import StoredPayLink +from lnbits.core.models.users import UserLabel +from lnbits.core.models.wallets import ( + Wallet, + WalletPermission, + WalletSharePermission, + WalletShareStatus, +) +from lnbits.db import dict_to_model + + +def test_user_label_uses_pydantic_v2_pattern_validation(): + label = UserLabel(name="label-1", color="#FF00AA") + + assert label.name == "label-1" + assert label.color == "#FF00AA" + + with pytest.raises(ValidationError): + UserLabel(name="label-1", color="bad-color") + + +def test_wallet_has_stored_paylinks_field_and_mirrors_it(): + source_wallet = Wallet( + id="source-wallet-id", + user="source-user-id", + name="source", + adminkey="admin-key", + inkey="invoice-key", + ) + source_wallet.stored_paylinks.links.append( + StoredPayLink(lnurl="lnurl1example", label="saved paylink") + ) + + shared_wallet = Wallet( + id="shared-wallet-id", + user="shared-user-id", + name="shared", + adminkey="shared-admin-key", + inkey="shared-invoice-key", + ) + source_wallet.extra.shared_with.append( + WalletSharePermission( + request_id="share-request-id", + username="shared-user", + shared_with_wallet_id=shared_wallet.id, + permissions=[WalletPermission.VIEW_PAYMENTS], + status=WalletShareStatus.APPROVED, + ) + ) + shared_wallet.mirror_shared_wallet(source_wallet) + + assert shared_wallet.stored_paylinks.links[0].lnurl == "lnurl1example" + assert shared_wallet.stored_paylinks.links[0].label == "saved paylink" + + +def test_db_dict_to_model_parses_optional_nested_pydantic_v2_models(): + ext = dict_to_model( + { + "id": "ext-id", + "name": "Extension", + "version": "1.0.0", + "active": 1, + "meta": ( + '{"installed_release": {"name": "Release", "version": "1.0.0", ' + '"archive": "https://example.com/release.zip", ' + '"source_repo": "lnbits/example"}, "payments": [], ' + '"dependencies": [], "featured": false, ' + '"has_paid_release": false, "has_free_release": false}' + ), + }, + InstallableExtension, + ) + + assert ext.meta is not None + assert ext.meta.installed_release is not None + assert ext.meta.installed_release.source_repo == "lnbits/example" diff --git a/tests/unit/test_settings.py b/tests/unit/test_settings.py index a300afb0f..7a43383c4 100644 --- a/tests/unit/test_settings.py +++ b/tests/unit/test_settings.py @@ -1,6 +1,7 @@ import pytest +from pydantic import ValidationError -from lnbits.settings import RedirectPath +from lnbits.settings import EditableSettings, RedirectPath, Settings, UpdateSettings lnurlp_redirect_path = { "from_path": "/.well-known/lnurlp", @@ -166,3 +167,60 @@ def test_redirect_path_new_path_from(lnurlp: RedirectPath): lnurlp.new_path_from("/.well-known/lnurlp/path/more") == "/lnurlp/api/v1/well-known/path/more" ) + + +def test_settings_loads_env_and_dotenv_values(tmp_path, monkeypatch): + env_file = tmp_path / ".env" + env_file.write_text( + "\n".join( + [ + "lnbits_admin_users=alice,bob", + "STRIKE_API_ENDPOINT=https://dotenv.strike.example", + ] + ), + encoding="utf-8", + ) + monkeypatch.chdir(tmp_path) + monkeypatch.setenv("LNBITS_ALLOWED_USERS", ' ["carol", "dave"] ') + monkeypatch.setenv("STRIKE_API_KEY", "secret-key") + monkeypatch.setenv("STRIKE_API_ENDPOINT", "https://env.strike.example") + + settings = Settings() + + assert settings.lnbits_admin_users == ["alice", "bob"] + assert settings.lnbits_allowed_users == ["carol", "dave"] + assert settings.strike_api_endpoint == "https://env.strike.example" + assert settings.strike_api_key == "secret-key" + + +def test_editable_settings_from_dict_ignores_unknown_fields(): + editable_settings = EditableSettings.from_dict( + {"lnbits_admin_users": ["alice"], "unknown_setting": "ignored"} + ) + + assert editable_settings.lnbits_admin_users == ["alice"] + assert "unknown_setting" not in editable_settings.model_dump() + + +def test_update_settings_splits_string_lists_and_forbids_extra_fields(): + updated = UpdateSettings.model_validate( + { + "lnbits_admin_users": "alice,bob", + "strike_api_endpoint": "https://api.strike.me/v1", + "strike_api_key": "secret-key", + } + ) + + assert updated.lnbits_admin_users == ["alice", "bob"] + assert updated.strike_api_endpoint == "https://api.strike.me/v1" + assert updated.strike_api_key == "secret-key" + + with pytest.raises(ValidationError): + UpdateSettings.model_validate({"unknown_setting": True}) + + +def test_update_settings_schema_does_not_include_env_names(): + schema = UpdateSettings.model_json_schema() + + assert "properties" in schema + assert all("env_names" not in prop for prop in schema["properties"].values()) diff --git a/uv.lock b/uv.lock index c7b4e5da6..161cf3fbb 100644 --- a/uv.lock +++ b/uv.lock @@ -1301,6 +1301,7 @@ dependencies = [ { name = "protobuf" }, { name = "pycryptodomex" }, { name = "pydantic" }, + { name = "pydantic-settings" }, { name = "pyjwt" }, { name = "pyln-client" }, { name = "pynostr" }, @@ -1385,6 +1386,7 @@ requires-dist = [ { name = "psycopg2-binary", marker = "extra == 'migration'", specifier = "~=2.9.11" }, { name = "pycryptodomex", specifier = "~=3.23.0" }, { name = "pydantic", specifier = "~=2.12.0" }, + { name = "pydantic-settings", specifier = ">=2.13.1" }, { name = "pyjwt", specifier = "~=2.12.0" }, { name = "pyln-client", specifier = "~=25.12.0" }, { name = "pynostr", specifier = "~=0.7.0" }, @@ -2095,6 +2097,20 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/36/c7/cfc8e811f061c841d7990b0201912c3556bfeb99cdcb7ed24adc8d6f8704/pydantic_core-2.41.5-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:56121965f7a4dc965bff783d70b907ddf3d57f6eba29b6d2e5dabfaf07799c51", size = 2145302, upload-time = "2025-11-04T13:43:46.64Z" }, ] +[[package]] +name = "pydantic-settings" +version = "2.13.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pydantic" }, + { name = "python-dotenv" }, + { name = "typing-inspection" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/52/6d/fffca34caecc4a3f97bda81b2098da5e8ab7efc9a66e819074a11955d87e/pydantic_settings-2.13.1.tar.gz", hash = "sha256:b4c11847b15237fb0171e1462bf540e294affb9b86db4d9aa5c01730bdbe4025", size = 223826, upload-time = "2026-02-19T13:45:08.055Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/00/4b/ccc026168948fec4f7555b9164c724cf4125eac006e176541483d2c959be/pydantic_settings-2.13.1-py3-none-any.whl", hash = "sha256:d56fd801823dbeae7f0975e1f8c8e25c258eb75d278ea7abb5d9cebb01b56237", size = 58929, upload-time = "2026-02-19T13:45:06.034Z" }, +] + [[package]] name = "pygments" version = "2.19.2"