diff --git a/.gitignore b/.gitignore index 34de0af..b6289f2 100644 --- a/.gitignore +++ b/.gitignore @@ -173,3 +173,6 @@ cython_debug/ # VSCode .vscode/ .devcontainer/ + +# PyCharm +.idea/ diff --git a/src/requela/builders/sqlalchemy.py b/src/requela/builders/sqlalchemy.py index 97c62bf..4cf25f1 100644 --- a/src/requela/builders/sqlalchemy.py +++ b/src/requela/builders/sqlalchemy.py @@ -41,7 +41,11 @@ def get_initial_query(self): return select(self.model_class) def get_field_type(self, field: str) -> type: - field_type = getattr(self.model_class, field).property.columns[0].type.python_type + model = self.model_class + parts = field.split(".") + for part in parts[:-1]: + model = getattr(model, part).property.mapper.class_ + field_type = getattr(model, parts[-1]).property.columns[0].type.python_type return field_type def apply_and(self, *conditions: ColumnExpressionArgument) -> ColumnElement: diff --git a/src/requela/rules.py b/src/requela/rules.py index e335c83..de90832 100644 --- a/src/requela/rules.py +++ b/src/requela/rules.py @@ -19,6 +19,7 @@ class FieldRule: allowed_operators: set[Operator] | None = None alias: str | None = None allow_ordering: bool = True + source: str | None = None @dataclass @@ -98,7 +99,9 @@ def _get_fields_documentation(self, alias_prefix: str = "") -> list[str]: @classmethod def _resolve_alias(cls, alias: str) -> str: - field_name, _ = cls._get_field_by_alias(alias) + field_name, field_def = cls._get_field_by_alias(alias) + if isinstance(field_def, FieldRule) and field_def.source: + return field_def.source return field_name @classmethod @@ -133,7 +136,8 @@ def _validate(self) -> None: for field_name, field_def in self._fields.items(): try: # Get field type from model (implementation depends on ORM) - field_type = self._get_field_type(field_name) + lookup = field_def.source or field_name + field_type = self._get_field_type(lookup) # If no operators specified, infer from field type if field_def.allowed_operators is None: @@ -200,7 +204,7 @@ def _get_relation_by_alias(cls, alias: str) -> tuple[str, RelationshipRule]: raise ValueError(f"Relation with alias '{alias}' not found") if len(relations) > 1: raise ValueError(f"Multiple relations found for alias '{alias}'") - return list(relations.items())[0] + return next(iter(relations.items())) @classmethod def _get_field_by_alias(cls, alias: str) -> tuple[str, FieldRule | RelationshipRule]: diff --git a/tests/sqlalchemy/rules.py b/tests/sqlalchemy/rules.py index a99167c..20bc25b 100644 --- a/tests/sqlalchemy/rules.py +++ b/tests/sqlalchemy/rules.py @@ -54,3 +54,6 @@ class UserRules(ModelRQLRules): account = RelationshipRule( rules=AccountRules(), ) + status = FieldRule( + source="account.status", + ) diff --git a/tests/sqlalchemy/test_rules.py b/tests/sqlalchemy/test_rules.py index 91e1957..331e4e1 100644 --- a/tests/sqlalchemy/test_rules.py +++ b/tests/sqlalchemy/test_rules.py @@ -7,7 +7,7 @@ from requela.dataclasses import Operator from requela.exceptions import RequelaError from requela.rules import FieldRule, ModelRQLRules -from tests.sqlalchemy.models import Account, Actor, User +from tests.sqlalchemy.models import Account, AccountStatus, Actor, User from tests.sqlalchemy.rules import AccountRules, UserRules from tests.sqlalchemy.utils import assert_statements_equal @@ -160,6 +160,7 @@ def test_get_documentation(): "|is_active|eq, ne|yes|", "|name|eq, ilike, in, like, ne, out|yes|", "|role|in, out|yes|", + "|status|eq, in, ne, out|yes|", ] @@ -189,3 +190,15 @@ def test_tenant_relationship_ne_invalid_value(): RequelaError, match="`ne` can be applied to relationship only to test for null." ): account_rules.build_query("ne(tenant,pip)") + + +def test_valid_source_field(): + user_rules = UserRules() + stmt = user_rules.build_query(f"eq(status,{AccountStatus.ACTIVE.value})") + account_alias = aliased(Account) + expected = ( + select(User) + .filter(account_alias.status == AccountStatus.ACTIVE) + .join(account_alias, User.account) + ) + assert_statements_equal(stmt, expected)