"""Business logic of the finance module (datamodel-finance.md §2 and §9). Every function filters by `user_id` and raises the French `AppError` subclasses of the core; routers stay thin (CONVENTIONS C2.6). """ import uuid from datetime import date from decimal import Decimal from typing import Any from sqlalchemy import Select, delete, func, or_, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from app.core.errors import ( ConflictError, DomainValidationError, ForbiddenError, NotFoundError, ) from app.modules.finance import stats as stats_service from app.modules.finance.categorize import validate_actions, validate_matchers from app.modules.finance.enums import ( AccountKind, CategoryKind, CategorySource, SourceKind, ) from app.modules.finance.models import ( FinAccount, FinBudget, FinCategory, FinImportRun, FinRule, FinSourceProfile, FinTransaction, ) from app.modules.finance.normalize import add_months, month_str from app.modules.finance.pipeline import transaction_dedup_hash from app.modules.finance.presets import ( BUILTIN_PROFILES, TRANSFER_CATEGORY_NAME, default_category_tree, ) from app.modules.imports.models import ImportRun MAX_DEPTH_MESSAGE = "L'arbre des catégories est limité à deux niveaux." SORTABLE_TRANSACTION_FIELDS = { "booked_date": FinTransaction.booked_date, "amount": FinTransaction.amount, "label_clean": FinTransaction.label_clean, } # --------------------------------------------------------------------------- # Lazy seeding (built-in profiles + default category tree) # --------------------------------------------------------------------------- def _commit_seed(db: Session) -> None: """Commit seed rows, tolerating a concurrent request that seeded first.""" try: db.commit() except IntegrityError: db.rollback() def ensure_builtin_profiles(db: Session) -> None: existing = set( db.scalars( select(FinSourceProfile.name).where(FinSourceProfile.user_id.is_(None)) ).all() ) created = False for preset in BUILTIN_PROFILES: if preset["name"] in existing: continue db.add( FinSourceProfile( id=uuid.uuid4(), user_id=None, name=preset["name"], kind=preset["kind"], config=preset["config"], is_builtin=True, ) ) created = True if created: _commit_seed(db) def ensure_user_categories(db: Session, user_id: int) -> None: exists = db.scalar( select(func.count()) .select_from(FinCategory) .where(FinCategory.user_id == user_id) ) if exists: return for root in default_category_tree(): parent = FinCategory( id=uuid.uuid4(), user_id=user_id, parent_id=None, name=root["name"], icon=root["icon"], color=root["color"], kind=root["kind"], is_system=root["is_system"], sort_order=root["sort_order"], ) db.add(parent) for child in root["children"]: db.add( FinCategory( id=uuid.uuid4(), user_id=user_id, parent_id=parent.id, name=child["name"], icon=None, color=None, kind=root["kind"], is_system=False, sort_order=child["sort_order"], ) ) _commit_seed(db) def ensure_seed(db: Session, user_id: int) -> None: ensure_builtin_profiles(db) ensure_user_categories(db, user_id) # --------------------------------------------------------------------------- # Accounts (§9.1) # --------------------------------------------------------------------------- def get_account(db: Session, user_id: int, account_id: uuid.UUID) -> FinAccount: account = db.get(FinAccount, account_id) if account is None or account.user_id != user_id: raise NotFoundError("Compte introuvable.") return account def _account_dict( account: FinAccount, balances: dict[uuid.UUID, Decimal], counts: dict[uuid.UUID, int], last_dates: dict[uuid.UUID, Any], ) -> dict[str, Any]: data = { column.name: getattr(account, column.name) for column in FinAccount.__table__.columns } data["balance"] = stats_service.money( stats_service.money(account.initial_balance) + balances.get(account.id, Decimal("0.00")) ) data["transaction_count"] = counts.get(account.id, 0) last = last_dates.get(account.id) data["last_transaction_date"] = ( date.fromisoformat(last) if isinstance(last, str) else last ) return data def _last_transaction_dates(db: Session, user_id: int) -> dict[uuid.UUID, Any]: return dict( db.execute( select(FinTransaction.account_id, func.max(FinTransaction.booked_date)) .where(FinTransaction.user_id == user_id) .group_by(FinTransaction.account_id) ).all() ) def list_accounts( db: Session, user_id: int, include_archived: bool = False ) -> list[dict[str, Any]]: stmt = select(FinAccount).where(FinAccount.user_id == user_id) if not include_archived: stmt = stmt.where(FinAccount.is_archived.is_(False)) accounts = list(db.scalars(stmt.order_by(FinAccount.name.asc())).all()) balances = stats_service.account_balances(db, user_id) counts = stats_service.account_transaction_counts(db, user_id) last_dates = _last_transaction_dates(db, user_id) return [ _account_dict(account, balances, counts, last_dates) for account in accounts ] def account_detail(db: Session, user_id: int, account_id: uuid.UUID) -> dict[str, Any]: account = get_account(db, user_id, account_id) return _account_dict( account, stats_service.account_balances(db, user_id), stats_service.account_transaction_counts(db, user_id), _last_transaction_dates(db, user_id), ) def create_account(db: Session, user_id: int, payload: Any) -> FinAccount: taken = db.scalar( select(func.count()) .select_from(FinAccount) .where(FinAccount.user_id == user_id, FinAccount.name == payload.name) ) if taken: raise ConflictError("Un compte porte déjà ce nom.") account = FinAccount( id=uuid.uuid4(), user_id=user_id, name=payload.name, kind=payload.kind, currency=payload.currency.upper(), institution=payload.institution, iban_masked=payload.iban_masked, initial_balance=payload.initial_balance, ) db.add(account) db.commit() return account def update_account( db: Session, user_id: int, account_id: uuid.UUID, payload: Any ) -> FinAccount: account = get_account(db, user_id, account_id) data = payload.model_dump(exclude_unset=True) if "name" in data and data["name"] != account.name: taken = db.scalar( select(func.count()) .select_from(FinAccount) .where( FinAccount.user_id == user_id, FinAccount.name == data["name"], FinAccount.id != account_id, ) ) if taken: raise ConflictError("Un compte porte déjà ce nom.") if data.get("currency"): data["currency"] = data["currency"].upper() for key, value in data.items(): setattr(account, key, value) db.commit() return account def delete_account(db: Session, user_id: int, account_id: uuid.UUID) -> None: account = get_account(db, user_id, account_id) count = db.scalar( select(func.count()) .select_from(FinTransaction) .where(FinTransaction.account_id == account.id) ) if count: raise ConflictError( "Ce compte contient des transactions. Archivez le compte à la place." ) db.delete(account) db.commit() # --------------------------------------------------------------------------- # Categories (§9.3) # --------------------------------------------------------------------------- def get_category(db: Session, user_id: int, category_id: uuid.UUID) -> FinCategory: category = db.get(FinCategory, category_id) if category is None or category.user_id != user_id: raise NotFoundError("Catégorie introuvable.") return category def category_tree(db: Session, user_id: int) -> list[dict[str, Any]]: categories = list( db.scalars( select(FinCategory) .where(FinCategory.user_id == user_id) .order_by(FinCategory.sort_order.asc(), FinCategory.name.asc()) ).all() ) counts = dict( db.execute( select(FinTransaction.category_id, func.count()) .where(FinTransaction.user_id == user_id) .group_by(FinTransaction.category_id) ).all() ) def to_dict(category: FinCategory) -> dict[str, Any]: return { "id": category.id, "parent_id": category.parent_id, "name": category.name, "icon": category.icon, "color": category.color, "kind": category.kind, "is_system": category.is_system, "sort_order": category.sort_order, "transaction_count": counts.get(category.id, 0), "children": [], } nodes = {c.id: to_dict(c) for c in categories} roots: list[dict[str, Any]] = [] for category in categories: node = nodes[category.id] if category.parent_id and category.parent_id in nodes: parent = nodes[category.parent_id] parent["children"].append(node) parent["transaction_count"] += node["transaction_count"] else: roots.append(node) return roots def create_category(db: Session, user_id: int, payload: Any) -> FinCategory: kind = payload.kind parent: FinCategory | None = None if payload.parent_id is not None: parent = get_category(db, user_id, payload.parent_id) if parent.parent_id is not None: raise DomainValidationError(MAX_DEPTH_MESSAGE) kind = parent.kind if kind == CategoryKind.TRANSFER: raise DomainValidationError( "La catégorie « Virements internes » est unique et gérée par l'application." ) sibling_clause = ( FinCategory.parent_id.is_(None) if payload.parent_id is None else FinCategory.parent_id == payload.parent_id ) taken = db.scalar( select(func.count()) .select_from(FinCategory) .where( FinCategory.user_id == user_id, FinCategory.name == payload.name, sibling_clause, ) ) if taken: raise ConflictError("Une catégorie porte déjà ce nom à ce niveau.") category = FinCategory( id=uuid.uuid4(), user_id=user_id, parent_id=payload.parent_id, name=payload.name, icon=payload.icon, color=payload.color, kind=kind, is_system=False, sort_order=payload.sort_order, ) db.add(category) db.commit() return category def update_category( db: Session, user_id: int, category_id: uuid.UUID, payload: Any ) -> FinCategory: category = get_category(db, user_id, category_id) data = payload.model_dump(exclude_unset=True) if category.is_system and "name" in data: raise DomainValidationError( "Cette catégorie système ne peut pas être renommée." ) if "parent_id" in data: new_parent_id = data["parent_id"] if new_parent_id is not None: has_children = db.scalar( select(func.count()) .select_from(FinCategory) .where(FinCategory.parent_id == category.id) ) if has_children: raise DomainValidationError(MAX_DEPTH_MESSAGE) parent = get_category(db, user_id, new_parent_id) if parent.parent_id is not None or parent.id == category.id: raise DomainValidationError(MAX_DEPTH_MESSAGE) category.kind = parent.kind for key, value in data.items(): setattr(category, key, value) db.commit() return category def delete_category(db: Session, user_id: int, category_id: uuid.UUID) -> None: category = get_category(db, user_id, category_id) if category.is_system: raise DomainValidationError( "Cette catégorie système ne peut pas être supprimée." ) child_ids = list( db.scalars( select(FinCategory.id).where(FinCategory.parent_id == category.id) ).all() ) all_ids = [category.id, *child_ids] db.execute( update(FinTransaction) .where( FinTransaction.user_id == user_id, FinTransaction.category_id.in_(all_ids), ) .values(category_id=None, category_source=None) ) db.execute( delete(FinBudget).where( FinBudget.user_id == user_id, FinBudget.category_id.in_(all_ids) ) ) # Rules pointing at a deleted category lose their action (disabled if empty). for rule in db.scalars(select(FinRule).where(FinRule.user_id == user_id)).all(): actions = dict(rule.actions or {}) target = actions.get("set_category_id") if target and uuid.UUID(str(target)) in set(all_ids): actions.pop("set_category_id", None) rule.actions = actions if not any( actions.get(key) for key in ("set_label_clean", "set_counterparty", "mark_transfer") ): rule.enabled = False db.delete(category) db.commit() # --------------------------------------------------------------------------- # Source profiles (§9.6) # --------------------------------------------------------------------------- def get_source_profile( db: Session, user_id: int, profile_id: uuid.UUID ) -> FinSourceProfile: profile = db.get(FinSourceProfile, profile_id) if profile is None or (profile.user_id is not None and profile.user_id != user_id): raise NotFoundError("Profil de source introuvable.") return profile def list_source_profiles(db: Session, user_id: int) -> list[FinSourceProfile]: stmt = ( select(FinSourceProfile) .where( or_( FinSourceProfile.user_id.is_(None), FinSourceProfile.user_id == user_id, ) ) .order_by(FinSourceProfile.is_builtin.desc(), FinSourceProfile.name.asc()) ) return list(db.scalars(stmt).all()) def find_profile_by_name(db: Session, name: str) -> FinSourceProfile | None: return db.scalars( select(FinSourceProfile).where( FinSourceProfile.user_id.is_(None), FinSourceProfile.name == name ) ).first() def _check_profile_name(db: Session, user_id: int, name: str) -> None: taken = db.scalar( select(func.count()) .select_from(FinSourceProfile) .where(FinSourceProfile.user_id == user_id, FinSourceProfile.name == name) ) if taken: raise ConflictError("Un profil de source porte déjà ce nom.") def create_source_profile(db: Session, user_id: int, payload: Any) -> FinSourceProfile: _check_profile_name(db, user_id, payload.name) profile = FinSourceProfile( id=uuid.uuid4(), user_id=user_id, name=payload.name, kind=payload.kind, config=payload.config, is_builtin=False, ) db.add(profile) db.commit() return profile def clone_source_profile( db: Session, user_id: int, profile_id: uuid.UUID ) -> FinSourceProfile: source = get_source_profile(db, user_id, profile_id) name = f"{source.name} (copie)" suffix = 2 while db.scalar( select(func.count()) .select_from(FinSourceProfile) .where(FinSourceProfile.user_id == user_id, FinSourceProfile.name == name) ): name = f"{source.name} (copie {suffix})" suffix += 1 clone = FinSourceProfile( id=uuid.uuid4(), user_id=user_id, name=name, kind=source.kind, config=dict(source.config or {}), is_builtin=False, ) db.add(clone) db.commit() return clone def update_source_profile( db: Session, user_id: int, profile_id: uuid.UUID, payload: Any ) -> FinSourceProfile: profile = get_source_profile(db, user_id, profile_id) if profile.is_builtin: raise ForbiddenError( "Les profils intégrés sont en lecture seule : dupliquez-le pour " "le modifier." ) data = payload.model_dump(exclude_unset=True) if "name" in data and data["name"] != profile.name: _check_profile_name(db, user_id, data["name"]) for key, value in data.items(): setattr(profile, key, value) db.commit() return profile def delete_source_profile(db: Session, user_id: int, profile_id: uuid.UUID) -> None: profile = get_source_profile(db, user_id, profile_id) if profile.is_builtin: raise ForbiddenError("Les profils intégrés ne peuvent pas être supprimés.") db.delete(profile) db.commit() # --------------------------------------------------------------------------- # Transactions (§9.2) # --------------------------------------------------------------------------- def transactions_query( db: Session, user_id: int, *, date_from: date | None = None, date_to: date | None = None, account_ids: list[uuid.UUID] | None = None, category_ids: list[str] | None = None, q: str | None = None, direction: str | None = None, amount_min: Decimal | None = None, amount_max: Decimal | None = None, is_transfer: bool | None = None, import_run_id: int | None = None, sort: str | None = None, ) -> Select: stmt = select(FinTransaction).where(FinTransaction.user_id == user_id) if date_from is not None: stmt = stmt.where(FinTransaction.booked_date >= date_from) if date_to is not None: stmt = stmt.where(FinTransaction.booked_date <= date_to) if account_ids: stmt = stmt.where(FinTransaction.account_id.in_(account_ids)) if category_ids: wants_none = any(str(c).lower() == "none" for c in category_ids) try: explicit = [ uuid.UUID(str(c)) for c in category_ids if str(c).lower() != "none" ] except ValueError as exc: raise DomainValidationError( "Filtre de catégorie invalide : identifiant ou « none » attendu." ) from exc expanded = set(explicit) if explicit: expanded |= set( db.scalars( select(FinCategory.id).where( FinCategory.user_id == user_id, FinCategory.parent_id.in_(explicit), ) ).all() ) clauses = [] if expanded: clauses.append(FinTransaction.category_id.in_(list(expanded))) if wants_none: clauses.append(FinTransaction.category_id.is_(None)) if clauses: stmt = stmt.where(or_(*clauses)) if q: pattern = f"%{q}%" stmt = stmt.where( or_( FinTransaction.label_clean.ilike(pattern), FinTransaction.label_raw.ilike(pattern), FinTransaction.counterparty.ilike(pattern), FinTransaction.notes.ilike(pattern), ) ) if direction == "debit": stmt = stmt.where(FinTransaction.amount < 0) elif direction == "credit": stmt = stmt.where(FinTransaction.amount > 0) if amount_min is not None: stmt = stmt.where(func.abs(FinTransaction.amount) >= amount_min) if amount_max is not None: stmt = stmt.where(func.abs(FinTransaction.amount) <= amount_max) if is_transfer is True: stmt = stmt.where(FinTransaction.transfer_group_id.is_not(None)) elif is_transfer is False: stmt = stmt.where(FinTransaction.transfer_group_id.is_(None)) if import_run_id is not None: stmt = stmt.where(FinTransaction.import_run_id == import_run_id) sort = sort or "-booked_date" descending = sort.startswith("-") column = SORTABLE_TRANSACTION_FIELDS.get(sort.lstrip("-")) if column is None: raise DomainValidationError( "Champ de tri non autorisé.", details={"sort": sort} ) order = column.desc() if descending else column.asc() return stmt.order_by(order, FinTransaction.id.desc()) def serialize_transactions( db: Session, user_id: int, transactions: list[FinTransaction] ) -> list[dict[str, Any]]: if not transactions: return [] accounts = { a.id: a for a in db.scalars( select(FinAccount).where(FinAccount.user_id == user_id) ).all() } categories = { c.id: c for c in db.scalars( select(FinCategory).where(FinCategory.user_id == user_id) ).all() } out: list[dict[str, Any]] = [] for tx in transactions: account = accounts.get(tx.account_id) category = categories.get(tx.category_id) if tx.category_id else None color = None if category is not None: color = category.color if color is None and category.parent_id: parent = categories.get(category.parent_id) color = parent.color if parent else None out.append( { "id": tx.id, "account_id": tx.account_id, "account_name": account.name if account else None, "booked_date": tx.booked_date, "value_date": tx.value_date, "amount": tx.amount, "currency": tx.currency, "label_raw": tx.label_raw, "label_clean": tx.label_clean, "counterparty": tx.counterparty, "category_id": tx.category_id, "category_name": category.name if category else None, "category_color": color, "category_source": tx.category_source, "notes": tx.notes, "transfer_group_id": tx.transfer_group_id, "external_id": tx.external_id, "import_run_id": tx.import_run_id, } ) return out def get_transaction( db: Session, user_id: int, transaction_id: uuid.UUID ) -> FinTransaction: tx = db.get(FinTransaction, transaction_id) if tx is None or tx.user_id != user_id: raise NotFoundError("Transaction introuvable.") return tx def create_transaction(db: Session, user_id: int, payload: Any) -> FinTransaction: account = get_account(db, user_id, payload.account_id) amount = Decimal(payload.amount).quantize(Decimal("0.01")) label = payload.label_clean.strip() dedup_hash = transaction_dedup_hash( db, account.id, payload.booked_date, amount, label ) category_source = CategorySource.USER if payload.category_id else None if payload.category_id: get_category(db, user_id, payload.category_id) tx = FinTransaction( id=uuid.uuid4(), user_id=user_id, account_id=account.id, booked_date=payload.booked_date, value_date=payload.value_date, amount=amount, currency=(payload.currency or account.currency or "EUR").upper(), label_raw=label, label_clean=label, counterparty=payload.counterparty, category_id=payload.category_id, category_source=category_source, notes=payload.notes, ) tx.dedup_hash = dedup_hash db.add(tx) db.commit() return tx def update_transaction( db: Session, user_id: int, transaction_id: uuid.UUID, payload: Any ) -> FinTransaction: tx = get_transaction(db, user_id, transaction_id) data = payload.model_dump(exclude_unset=True) if ("amount" in data or "booked_date" in data) and tx.import_run_id is not None: raise DomainValidationError( "La date et le montant d'une transaction importée ne sont pas modifiables." ) if "category_id" in data: if data["category_id"] is None: tx.category_id = None tx.category_source = None else: get_category(db, user_id, data["category_id"]) tx.category_id = data["category_id"] tx.category_source = CategorySource.USER data.pop("category_id") for key, value in data.items(): setattr(tx, key, value) db.commit() return tx def delete_transaction(db: Session, user_id: int, transaction_id: uuid.UUID) -> None: tx = get_transaction(db, user_id, transaction_id) if tx.import_run_id is not None: raise DomainValidationError( "Cette transaction provient d'un import : supprimez l'import " "complet ou ignorez la ligne." ) db.delete(tx) db.commit() def bulk_categorize( db: Session, user_id: int, transaction_ids: list[uuid.UUID], category_id: uuid.UUID | None, ) -> int: if category_id is not None: get_category(db, user_id, category_id) result = db.execute( update(FinTransaction) .where( FinTransaction.user_id == user_id, FinTransaction.id.in_(transaction_ids), ) .values( category_id=category_id, category_source=CategorySource.USER if category_id else None, ) ) db.commit() return int(result.rowcount or 0) # --------------------------------------------------------------------------- # Rules (§9.4) # --------------------------------------------------------------------------- def get_rule(db: Session, user_id: int, rule_id: uuid.UUID) -> FinRule: rule = db.get(FinRule, rule_id) if rule is None or rule.user_id != user_id: raise NotFoundError("Règle introuvable.") return rule def list_rules(db: Session, user_id: int) -> list[FinRule]: stmt = ( select(FinRule) .where(FinRule.user_id == user_id) .order_by(FinRule.priority.asc(), FinRule.created_at.asc(), FinRule.id.asc()) ) return list(db.scalars(stmt).all()) def _validate_rule_payload( db: Session, user_id: int, matchers: dict[str, Any], actions: dict[str, Any] ) -> None: validate_matchers(matchers) validate_actions(actions) category_id = actions.get("set_category_id") if category_id: get_category(db, user_id, uuid.UUID(str(category_id))) account_id = matchers.get("account_id") if account_id: get_account(db, user_id, uuid.UUID(str(account_id))) def create_rule(db: Session, user_id: int, payload: Any) -> FinRule: _validate_rule_payload(db, user_id, payload.matchers, payload.actions) rule = FinRule( id=uuid.uuid4(), user_id=user_id, name=payload.name, priority=payload.priority, enabled=payload.enabled, stop=payload.stop, matchers=payload.matchers, actions=payload.actions, ) db.add(rule) db.commit() return rule def update_rule(db: Session, user_id: int, rule_id: uuid.UUID, payload: Any) -> FinRule: rule = get_rule(db, user_id, rule_id) data = payload.model_dump(exclude_unset=True) matchers = data.get("matchers", rule.matchers) actions = data.get("actions", rule.actions) _validate_rule_payload(db, user_id, matchers, actions) for key, value in data.items(): setattr(rule, key, value) db.commit() return rule def delete_rule(db: Session, user_id: int, rule_id: uuid.UUID) -> None: rule = get_rule(db, user_id, rule_id) db.delete(rule) db.commit() def reorder_rules(db: Session, user_id: int, ordered_ids: list[uuid.UUID]) -> None: rules = {rule.id: rule for rule in list_rules(db, user_id)} for index, rule_id in enumerate(ordered_ids): rule = rules.get(rule_id) if rule is None: raise NotFoundError("Règle introuvable.") rule.priority = index * 10 db.commit() # --------------------------------------------------------------------------- # Budgets (§9.5) # --------------------------------------------------------------------------- def get_budget(db: Session, user_id: int, budget_id: uuid.UUID) -> FinBudget: budget = db.get(FinBudget, budget_id) if budget is None or budget.user_id != user_id: raise NotFoundError("Budget introuvable.") return budget def _check_budget_overlap( db: Session, user_id: int, category_id: uuid.UUID, start_month: date, end_month: date | None, exclude_id: uuid.UUID | None = None, ) -> None: stmt = select(FinBudget).where( FinBudget.user_id == user_id, FinBudget.category_id == category_id ) if exclude_id is not None: stmt = stmt.where(FinBudget.id != exclude_id) for other in db.scalars(stmt).all(): starts_after_other_ends = ( other.end_month is not None and start_month > other.end_month ) ends_before_other_starts = ( end_month is not None and end_month < other.start_month ) if not starts_after_other_ends and not ends_before_other_starts: raise ConflictError( "Un budget existe déjà pour cette catégorie sur cette période." ) def list_budgets( db: Session, user_id: int, month: date, account_ids: list[uuid.UUID] | None = None ) -> list[dict[str, Any]]: progress = stats_service.budget_progress( db, user_id, month, account_ids=account_ids ) by_id = {item["budget_id"]: item for item in progress["items"]} budgets = db.scalars( select(FinBudget) .where( FinBudget.user_id == user_id, FinBudget.start_month <= month.replace(day=1), or_( FinBudget.end_month.is_(None), FinBudget.end_month >= month.replace(day=1), ), ) .order_by(FinBudget.start_month.asc()) ).all() out: list[dict[str, Any]] = [] for budget in budgets: item = by_id.get(budget.id, {}) out.append( { "id": budget.id, "category_id": budget.category_id, "category_name": item.get("category_name"), "category_color": item.get("category_color"), "monthly_amount": budget.monthly_amount, "start_month": budget.start_month, "end_month": budget.end_month, "actual": item.get("actual", Decimal("0.00")), "remaining": item.get( "remaining", stats_service.money(budget.monthly_amount) ), "progress_pct": item.get("progress_pct", 0.0), "projected_eom": item.get("projected_eom"), "status": item.get("status", "ok"), } ) return out def create_budget(db: Session, user_id: int, payload: Any) -> FinBudget: category = get_category(db, user_id, payload.category_id) if category.kind != CategoryKind.EXPENSE: raise DomainValidationError( "Un budget ne peut cibler qu'une catégorie de dépense." ) if payload.end_month is not None and payload.end_month < payload.start_month: raise DomainValidationError( "Le mois de fin doit être postérieur au mois de début." ) _check_budget_overlap( db, user_id, payload.category_id, payload.start_month, payload.end_month ) budget = FinBudget( id=uuid.uuid4(), user_id=user_id, category_id=payload.category_id, monthly_amount=payload.monthly_amount, start_month=payload.start_month, end_month=payload.end_month, ) db.add(budget) db.commit() return budget def update_budget( db: Session, user_id: int, budget_id: uuid.UUID, payload: Any ) -> list[FinBudget]: budget = get_budget(db, user_id, budget_id) data = payload.model_dump(exclude_unset=True) effective_from = data.pop("effective_from", None) if effective_from is not None: # Close the current budget and open a new one from `effective_from`. amount = data.get("monthly_amount", budget.monthly_amount) if effective_from <= budget.start_month: budget.monthly_amount = amount db.commit() return [budget] budget.end_month = add_months(effective_from, -1) successor = FinBudget( id=uuid.uuid4(), user_id=user_id, category_id=budget.category_id, monthly_amount=amount, start_month=effective_from, end_month=data.get("end_month"), ) db.add(successor) db.commit() return [budget, successor] if "end_month" in data: _check_budget_overlap( db, user_id, budget.category_id, budget.start_month, data["end_month"], exclude_id=budget.id, ) for key, value in data.items(): setattr(budget, key, value) db.commit() return [budget] def delete_budget(db: Session, user_id: int, budget_id: uuid.UUID) -> None: budget = get_budget(db, user_id, budget_id) db.delete(budget) db.commit() # --------------------------------------------------------------------------- # Import runs (§9.6) # --------------------------------------------------------------------------- def import_runs_query(user_id: int) -> Select: return ( select(FinImportRun) .where(FinImportRun.user_id == user_id) .order_by(FinImportRun.started_at.desc(), FinImportRun.id.desc()) ) def get_import_run(db: Session, user_id: int, run_id: uuid.UUID) -> FinImportRun: run = db.get(FinImportRun, run_id) if run is None or run.user_id != user_id: raise NotFoundError("Import introuvable.") return run def rollback_import( db: Session, user_id: int, run_id: uuid.UUID, force: bool = False ) -> None: """Delete the central run: transactions vanish through ON DELETE CASCADE.""" run = get_import_run(db, user_id, run_id) if not force: touched = db.scalar( select(func.count()) .select_from(FinTransaction) .where( FinTransaction.import_run_id == run.import_run_id, or_( FinTransaction.category_source == CategorySource.USER, FinTransaction.notes.is_not(None), ), ) ) if touched: raise ConflictError( "Des transactions de cet import ont été modifiées " "manuellement. Relancez avec force=true pour les supprimer " "malgré tout." ) central = db.get(ImportRun, run.import_run_id) if central is not None: db.delete(central) else: db.delete(run) db.commit() def resolve_profile( db: Session, user_id: int, profile_id: uuid.UUID | None ) -> FinSourceProfile: if profile_id is not None: return get_source_profile(db, user_id, profile_id) profile = find_profile_by_name(db, "CSV générique") if profile is None: raise NotFoundError("Profil de source introuvable.") return profile def current_month(today: date) -> date: return today.replace(day=1) def parse_month_param(value: str | None, today: date) -> date: if not value: return current_month(today) try: year, month = value.split("-") return date(int(year), int(month), 1) except (ValueError, TypeError) as exc: raise DomainValidationError( f"Mois invalide : « {value} » (format attendu AAAA-MM)." ) from exc def transfer_category(db: Session, user_id: int) -> FinCategory | None: return db.scalars( select(FinCategory).where( FinCategory.user_id == user_id, FinCategory.name == TRANSFER_CATEGORY_NAME, ) ).first() def account_for_source_kind( db: Session, user_id: int, kind: str, label: str ) -> FinAccount: """Default target account used by the generic `POST /api/imports` route.""" account = db.scalars( select(FinAccount).where( FinAccount.user_id == user_id, FinAccount.name == label ) ).first() if account is not None: return account account = FinAccount( id=uuid.uuid4(), user_id=user_id, name=label, kind=AccountKind.PAYPAL if kind == SourceKind.PAYPAL_CSV else AccountKind.CHECKING, currency="EUR", institution=None, initial_balance=Decimal("0.00"), ) db.add(account) db.flush() return account def month_label(value: date) -> str: return month_str(value)