diff --git a/python/aioia_core/__init__.py b/python/aioia_core/__init__.py index b74b5ba..89856ce 100644 --- a/python/aioia_core/__init__.py +++ b/python/aioia_core/__init__.py @@ -20,9 +20,16 @@ ) from aioia_core.factories.base_repository_factory import BaseRepositoryFactory from aioia_core.models import Base, BaseModel -from aioia_core.protocols import ( +from aioia_core.types import ( + ConditionalFilter, + ConditionalOperator, + CrudFilter, CrudRepositoryProtocol, DatabaseRepositoryProtocol, + FilterOperator, + LogicalFilter, + is_conditional_filter, + is_logical_filter, ) from aioia_core.repositories import BaseRepository from aioia_core.settings import DatabaseSettings, JWTSettings, OpenAIAPISettings @@ -30,7 +37,7 @@ # Deprecated imports for backwards compatibility from aioia_core.factories.base_manager_factory import BaseManagerFactory from aioia_core.managers import BaseManager -from aioia_core.protocols import CrudManagerProtocol, DatabaseManagerProtocol +from aioia_core.types import CrudManagerProtocol, DatabaseManagerProtocol __all__ = [ # Database - New names (recommended) @@ -53,6 +60,14 @@ "INTERNAL_SERVER_ERROR", "extract_error_code_from_exception", "get_error_detail_from_exception", + # Filters + "CrudFilter", + "LogicalFilter", + "ConditionalFilter", + "FilterOperator", + "ConditionalOperator", + "is_logical_filter", + "is_conditional_filter", # Settings "DatabaseSettings", "OpenAIAPISettings", diff --git a/python/aioia_core/factories/base_manager_factory.py b/python/aioia_core/factories/base_manager_factory.py index db080fc..34962a5 100644 --- a/python/aioia_core/factories/base_manager_factory.py +++ b/python/aioia_core/factories/base_manager_factory.py @@ -20,7 +20,7 @@ # Re-export from base_repository_factory module from aioia_core.factories.base_repository_factory import BaseRepositoryFactory -from aioia_core.protocols import DatabaseRepositoryProtocol +from aioia_core.types import DatabaseRepositoryProtocol # TypeVar for backwards compatibility (cannot alias TypeVar directly) ManagerType = TypeVar("ManagerType", bound=DatabaseRepositoryProtocol) diff --git a/python/aioia_core/factories/base_repository_factory.py b/python/aioia_core/factories/base_repository_factory.py index c5e910e..f773684 100644 --- a/python/aioia_core/factories/base_repository_factory.py +++ b/python/aioia_core/factories/base_repository_factory.py @@ -6,7 +6,7 @@ from sqlalchemy.orm import Session, sessionmaker -from aioia_core.protocols import DatabaseRepositoryProtocol +from aioia_core.types import DatabaseRepositoryProtocol # RepositoryType을 DatabaseRepositoryProtocol에 바인딩 RepositoryType = TypeVar("RepositoryType", bound=DatabaseRepositoryProtocol) diff --git a/python/aioia_core/fastapi/base_crud_router.py b/python/aioia_core/fastapi/base_crud_router.py index a3a2a28..a3c56e7 100644 --- a/python/aioia_core/fastapi/base_crud_router.py +++ b/python/aioia_core/fastapi/base_crud_router.py @@ -2,7 +2,7 @@ import warnings from collections.abc import Callable, Sequence from datetime import datetime, timezone -from typing import Any, Generic, TypeVar +from typing import Any, Generic, TypeVar, cast import sentry_sdk from fastapi import APIRouter, Body, Depends, HTTPException, Query, status @@ -23,7 +23,14 @@ RESOURCE_UPDATE_FAILED, ErrorResponse, ) -from aioia_core.protocols import DatabaseRepositoryProtocol, ModelType, RepositoryType +from aioia_core.types import ( + CrudFilter, + DatabaseRepositoryProtocol, + ModelType, + RepositoryType, + is_conditional_filter, + is_logical_filter, +) # TypeVar for _create_repository_dependency_from_factory method FactoryRepositoryType = TypeVar("FactoryRepositoryType", bound=DatabaseRepositoryProtocol) @@ -536,36 +543,29 @@ def _get_item_or_404(self, repository: RepositoryType, item_id: str) -> ModelTyp return item def _decamelize_filter_fields( - self, filters: list[dict[str, Any]] - ) -> list[dict[str, Any]]: + self, filters: list[CrudFilter] + ) -> list[CrudFilter]: """Recursively traverses the filter structure and decamelizes field names.""" - processed_filters = [] + processed_filters: list[Any] = [] for filter_item in filters: - # Conditional Filter (or/and) - if ( - filter_item.get("operator") in {"or", "and"} - and "value" in filter_item - and isinstance(filter_item["value"], list) - ): + if is_conditional_filter(filter_item): processed_filters.append( { **filter_item, "value": self._decamelize_filter_fields(filter_item["value"]), } ) - # Logical Filter - elif "field" in filter_item: + elif is_logical_filter(filter_item): processed_filters.append( { **filter_item, - "field": decamelize(str(filter_item["field"])), + "field": decamelize(filter_item["field"]), } ) - # Unrecognized filter structure, append as is else: processed_filters.append(filter_item) - return processed_filters + return cast(list[CrudFilter], processed_filters) def _parse_query_params( self, sort_param: str | None, filters_param: str | None diff --git a/python/aioia_core/repositories.py b/python/aioia_core/repositories.py index 722e45a..f8526e8 100644 --- a/python/aioia_core/repositories.py +++ b/python/aioia_core/repositories.py @@ -13,10 +13,11 @@ from uuid import uuid4 from pydantic import BaseModel as PydanticBaseModel -from sqlalchemy import and_, desc, or_ +from sqlalchemy import ColumnElement, and_, desc, or_ from sqlalchemy.orm import Session from aioia_core.models import BaseModel +from aioia_core.types import CrudFilter, is_conditional_filter, is_logical_filter ModelType = TypeVar("ModelType", bound=PydanticBaseModel) DBModelType = TypeVar("DBModelType", bound=BaseModel) @@ -76,7 +77,7 @@ def get_all( current: int = 1, page_size: int = 10, sort: list[tuple[str, str]] | None = None, - filters: list[dict[str, Any]] | None = None, + filters: list[CrudFilter] | None = None, load_options: list[Any] | None = None, ) -> tuple[list[ModelType], int]: """ @@ -134,31 +135,27 @@ def get_all( return [self.convert_to_model(item) for item in db_items], total - def _build_filter_conditions(self, filters: list[dict[str, Any]]) -> list[Any]: + def _build_filter_conditions( + self, filters: list[CrudFilter] + ) -> list[ColumnElement[bool]]: """Recursively builds SQLAlchemy filter conditions from filter criteria.""" - conditions = [] + conditions: list[ColumnElement[bool]] = [] for filter_item in filters: - operator = filter_item.get("operator") - - # Conditional Filter (or/and) - if ( - operator in {"or", "and"} - and "value" in filter_item - and isinstance(filter_item["value"], list) - ): + if is_conditional_filter(filter_item): nested_conditions = self._build_filter_conditions(filter_item["value"]) if nested_conditions: - if operator == "or": + if filter_item["operator"] == "or": conditions.append(or_(*nested_conditions)) else: conditions.append(and_(*nested_conditions)) continue - # Logical Filter - field = filter_item.get("field") - if not field: + if not is_logical_filter(filter_item): continue + field = filter_item["field"] + operator = filter_item["operator"] + column = getattr(self.db_model, field, None) if column is None: continue diff --git a/python/aioia_core/testing/crud_fixtures.py b/python/aioia_core/testing/crud_fixtures.py index 716be12..3aaae49 100644 --- a/python/aioia_core/testing/crud_fixtures.py +++ b/python/aioia_core/testing/crud_fixtures.py @@ -6,7 +6,7 @@ from sqlalchemy import DateTime, Integer, String, or_ from sqlalchemy.orm import DeclarativeBase, Mapped, Session, mapped_column -from aioia_core.protocols import DatabaseRepositoryProtocol +from aioia_core.types import DatabaseRepositoryProtocol from aioia_core.factories.base_repository_factory import BaseRepositoryFactory @@ -63,24 +63,30 @@ def get_all(self, current=1, page_size=10, sort=None, filters=None): if filters: # Simplified filter handling for tests, not a full implementation for f in filters: - if f["operator"] == "eq": - column = getattr(TestDBModel, f["field"]) - value = f["value"] + op = f.get("operator") + field = f.get("field") + value = f.get("value") + + if op == "eq" and field: + column = getattr(TestDBModel, field) if isinstance(column.type, DateTime) and isinstance(value, str): value = datetime.fromisoformat(value) q = q.filter(column == value) - elif f["operator"] == "in": - q = q.filter(getattr(TestDBModel, f["field"]).in_(f["value"])) - elif f["operator"] == "null": - q = q.filter(getattr(TestDBModel, f["field"]).is_(None)) - elif f["operator"] == "nnull": - q = q.filter(getattr(TestDBModel, f["field"]).isnot(None)) - elif f["operator"] == "or": + elif op == "in" and field: + q = q.filter(getattr(TestDBModel, field).in_(value)) + elif op == "null" and field: + q = q.filter(getattr(TestDBModel, field).is_(None)) + elif op == "nnull" and field: + q = q.filter(getattr(TestDBModel, field).isnot(None)) + elif op == "or" and isinstance(value, list): or_conditions = [] - for or_f in f["value"]: - if or_f["operator"] == "eq": + for or_f in value: + or_op = or_f.get("operator") + or_field = or_f.get("field") + or_value = or_f.get("value") + if or_op == "eq" and or_field: or_conditions.append( - getattr(TestDBModel, or_f["field"]) == or_f["value"] + getattr(TestDBModel, or_field) == or_value ) q = q.filter(or_(*or_conditions)) diff --git a/python/aioia_core/protocols.py b/python/aioia_core/types.py similarity index 69% rename from python/aioia_core/protocols.py rename to python/aioia_core/types.py index df73649..902efaa 100644 --- a/python/aioia_core/protocols.py +++ b/python/aioia_core/types.py @@ -1,16 +1,86 @@ """ -CRUD repository protocol definition for AIoIA projects. +CRUD repository protocol and type definitions for AIoIA projects. -Defines the interface for generic CRUD operations. +Defines the interface for generic CRUD operations and filter types. """ from __future__ import annotations -from typing import Any, Generic, Protocol, TypeVar +from typing import ( + Any, + Generic, + Literal, + NotRequired, + Protocol, + TypedDict, + TypeGuard, + TypeVar, +) from pydantic import BaseModel from sqlalchemy.orm import Session +# Filter type definitions (compatible with Refine's filter structure) +FilterOperator = Literal[ + "eq", + "ne", + "gt", + "gte", + "lt", + "lte", + "in", + "contains", + "startswith", + "endswith", + "null", + "nnull", +] + +ConditionalOperator = Literal["or", "and"] + + +class LogicalFilter(TypedDict): + """ + Single field filter condition. + + Example: + {"field": "status", "operator": "eq", "value": "active"} + {"field": "status", "operator": "null"} # value not required for null/nnull + """ + + field: str + operator: FilterOperator + value: NotRequired[Any] + + +class ConditionalFilter(TypedDict): + """ + OR/AND combination filter. + + Example: + {"operator": "or", "value": [ + {"field": "status", "operator": "eq", "value": "active"}, + {"field": "status", "operator": "eq", "value": "pending"} + ]} + """ + + operator: ConditionalOperator + value: list[CrudFilter] + + +CrudFilter = LogicalFilter | ConditionalFilter + + +def is_logical_filter(f: CrudFilter) -> TypeGuard[LogicalFilter]: + """Type guard to narrow CrudFilter to LogicalFilter.""" + return "field" in f + + +def is_conditional_filter(f: CrudFilter) -> TypeGuard[ConditionalFilter]: + """Type guard to narrow CrudFilter to ConditionalFilter.""" + return "field" not in f and "operator" in f + + ModelType = TypeVar("ModelType", bound=BaseModel) CreateSchemaType_contra = TypeVar( "CreateSchemaType_contra", bound=BaseModel, contravariant=True @@ -51,7 +121,7 @@ def get_all( current: int = 1, page_size: int = 10, sort: list[tuple[str, str]] | None = None, - filters: list[dict[str, Any]] | None = None, + filters: list[CrudFilter] | None = None, ) -> tuple[list[ModelType], int]: """ Retrieve all items with pagination, sorting, and filtering. @@ -61,7 +131,7 @@ def get_all( page_size: Number of items per page sort: Sort criteria as [(field, order), ...] where order is 'asc' or 'desc' Example: [('created_at', 'desc'), ('name', 'asc')] - filters: Filter conditions as [{'field': str, 'operator': str, 'value': Any}, ...] + filters: Filter conditions as list of CrudFilter (LogicalFilter or ConditionalFilter) Supported operators: eq, ne, contains, gt, gte, lt, lte, in, null, nnull, or, and Example: [{'field': 'status', 'operator': 'eq', 'value': 'active'}] @@ -140,16 +210,3 @@ def __init__(self, db_session: Session) -> None: # TypeVar aliases need to be redefined (cannot alias TypeVar directly) ManagerType = TypeVar("ManagerType", bound=CrudRepositoryProtocol) - -# For re-export compatibility, also export ModelType -__all__ = [ - # New names (recommended) - "CrudRepositoryProtocol", - "DatabaseRepositoryProtocol", - "RepositoryType", - "ModelType", - # Deprecated aliases (backwards compatibility) - "CrudManagerProtocol", - "DatabaseManagerProtocol", - "ManagerType", -] diff --git a/python/tests/unit/test_base_repository.py b/python/tests/unit/test_base_repository.py index e3ccde0..9dfedf2 100644 --- a/python/tests/unit/test_base_repository.py +++ b/python/tests/unit/test_base_repository.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Mapped, mapped_column from aioia_core.repositories import BaseRepository +from aioia_core.types import CrudFilter from aioia_core.models import BaseModel as DBBaseModel from aioia_core.testing.database_manager import TestDatabaseManager @@ -280,7 +281,7 @@ def test_conditional_filters(self): # OR 조건 테스트 # title이 'Apple' 이거나 content가 'Yellow fruit'인 경우 - filters = [ + or_filters: list[CrudFilter] = [ { "operator": "or", "value": [ @@ -289,7 +290,7 @@ def test_conditional_filters(self): ], } ] - items, total = self.repository.get_all(filters=filters) + items, total = self.repository.get_all(filters=or_filters) self.assertEqual(total, 3) titles = {item.title for item in items} self.assertIn("Apple", titles) @@ -297,7 +298,7 @@ def test_conditional_filters(self): # AND와 OR 중첩 조건 테스트 # (title이 'Apple' AND content가 'Red fruit') OR (title이 'Banana') - filters = [ + nested_filters: list[CrudFilter] = [ { "operator": "or", "value": [ @@ -316,7 +317,7 @@ def test_conditional_filters(self): ], } ] - items, total = self.repository.get_all(filters=filters) + items, total = self.repository.get_all(filters=nested_filters) self.assertEqual(total, 2) retrieved_titles = {item.title for item in items} self.assertEqual(retrieved_titles, {"Apple", "Banana"})