diff --git a/examples/conformance/models/platform.py b/examples/conformance/models/platform.py index 2fc10eb..33aaa40 100644 --- a/examples/conformance/models/platform.py +++ b/examples/conformance/models/platform.py @@ -9,6 +9,16 @@ class Platform: def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() if "id" in kwargs and "id" not in kwargs: diff --git a/examples/conformance/models/work_item.py b/examples/conformance/models/work_item.py index d5199f9..d3532c6 100644 --- a/examples/conformance/models/work_item.py +++ b/examples/conformance/models/work_item.py @@ -10,6 +10,16 @@ class WorkItem: def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() if "id" in kwargs and "id" not in kwargs: diff --git a/examples/conformance/pyproject.toml b/examples/conformance/pyproject.toml index ba426ff..182c898 100644 --- a/examples/conformance/pyproject.toml +++ b/examples/conformance/pyproject.toml @@ -2,14 +2,14 @@ name = "runtime-example-conformance-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] [tool.setuptools.packages.find] where = ["."] -include = ["models*", "requests*", "teaql*"] +include = ["models*", "requests*"] [build-system] requires = ["setuptools>=42"] diff --git a/examples/conformance/runtime_module.py b/examples/conformance/runtime_module.py index 568c7f7..0ccf55c 100644 --- a/examples/conformance/runtime_module.py +++ b/examples/conformance/runtime_module.py @@ -1,5 +1,9 @@ +import asyncio from datetime import datetime, timezone -from teaql.runtime import CheckResult, ObjectLocation, RuntimeModule +from teaql.runtime import CheckResult, ContextEntityRef, JsonFieldNamingProfile, ObjectLocation, RuntimeModule, create_wire_entity_metadata +from teaql.core.meta import EntityDescriptor, PropertyDescriptor, RelationDescriptor +from teaql.core.value import DataType +from Q import Q from teaql.core.value import Value try: from teaql.core.graph import GraphNode @@ -57,10 +61,60 @@ def check_and_fix(self, context, record, location, results): +_Platform_DESCRIPTOR = (EntityDescriptor("Platform") + .table_name("platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("work_item_list", "WorkItem").local("id").foreign("platform").many()) +) + +_WorkItem_DESCRIPTOR = (EntityDescriptor("WorkItem") + .table_name("work_item_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("title", DataType.Text).column_name("title").required()).property(PropertyDescriptor("description", DataType.Text).column_name("description")).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")) +) + +async def _ensure_generated_bootstrap_once(context): + previous_actor = context.user_identifier() if hasattr(context, 'user_identifier') else None + previous_category = context.get_resource('bootstrapCategory') + if hasattr(context, 'set_user_identifier'): + context.set_user_identifier('teaql-generated-bootstrap') + context.insert_resource('bootstrapCategory', 'runtime-bootstrap') + try: + platform_1 = await (Q.platforms().with_id_is(1).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if platform_1 is None: + platform_1 = Platform._teaql_new_with_fixed_id(1) + platform_1.update_name("Runtime Example") + try: + await platform_1.audit_as('create model root Platform(1)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + platform_1 = await (Q.platforms().with_id_is(1).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if platform_1 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if platform_1 is None: + raise _teaql_create_error + context.with_active_root(ContextEntityRef("Platform", 1)) + finally: + if hasattr(context, 'set_user_identifier'): + context.set_user_identifier(previous_actor) + context.insert_resource('bootstrapCategory', previous_category) + +async def _ensure_generated_bootstrap(context): + for _teaql_attempt in range(5): + try: + await _ensure_generated_bootstrap_once(context) + return + except Exception: + if _teaql_attempt == 4: + raise + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + + # Passive generated manifest. Call ensure_schema() separately and explicitly. GENERATED_RUNTIME_MODULE = (RuntimeModule().entity(Platform) - .checker("Platform", _PlatformChecker()).entity(WorkItem) + .schema_entity(_Platform_DESCRIPTOR) + .checker("Platform", _PlatformChecker()) + .wire_metadata("Platform", create_wire_entity_metadata("Platform", ["id", "name", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "version": ["version"]})).entity(WorkItem) + .schema_entity(_WorkItem_DESCRIPTOR) .checker("WorkItem", _WorkItemChecker()) - - .root_graph(GraphNode("Platform").set("id", 1).set("name", "Runtime Example")) + .wire_metadata("WorkItem", create_wire_entity_metadata("WorkItem", ["id", "title", "description", "platform", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "title": ["title"], "description": ["description"], "platform": ["platform"], "version": ["version"]})) + .generated_bootstrap(_ensure_generated_bootstrap) ) \ No newline at end of file diff --git a/examples/conformance/teaql/__init__.py b/examples/conformance/teaql/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/conformance/teaql/core/__init__.py b/examples/conformance/teaql/core/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/conformance/teaql/core/expr.py b/examples/conformance/teaql/core/expr.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/core/expr.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/core/list.py b/examples/conformance/teaql/core/list.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/core/list.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/core/mutation.py b/examples/conformance/teaql/core/mutation.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/core/mutation.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/core/query.py b/examples/conformance/teaql/core/query.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/core/query.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/core/value.py b/examples/conformance/teaql/core/value.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/core/value.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/data_service.py b/examples/conformance/teaql/data_service.py deleted file mode 100644 index b166884..0000000 --- a/examples/conformance/teaql/data_service.py +++ /dev/null @@ -1,1375 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "version": "integer"}, - "required": {"id": True, "name": True, "version": True}, - "relations": {**{}, **{"work_item_list": {"target_entity": "WorkItem", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"WorkItem": { - "table": "work_item_data", - "columns": {"id": "integer", "title": "text", "description": "text", "platform": "integer", "version": "integer"}, - "required": {"id": True, "title": True, "description": False, "platform": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/conformance/teaql/runtime.py b/examples/conformance/teaql/runtime.py deleted file mode 100644 index 904082f..0000000 --- a/examples/conformance/teaql/runtime.py +++ /dev/null @@ -1,619 +0,0 @@ -from dataclasses import dataclass -from datetime import timedelta -from enum import Enum -import asyncio -import builtins -import contextvars -import time - -_ID_SET_STORE = {} -_ID_SET_LOCKS = {} - -_SCHEMA_INVOCATION = object() - -_CHECK_MESSAGES = { - "en": { - "required": "{location} is required", - "min": "{location} is below the minimum", - "max": "{location} exceeds the maximum", - "min_length": "{location} is too short", - "max_length": "{location} is too long", - }, - "zh-CN": { - "required": "{location} 为必填项", - "min": "{location} 小于最小值", - "max": "{location} 超过最大值", - "min_length": "{location} 长度不足", - "max_length": "{location} 长度过长", - }, -} -_SUPPORTED_LOCALES = {"en", "zh-CN", "zh-TW", "ja", "ko", "de", "fr", "es", "pt", "ar", "th", "id", "fil", "uk", "vi"} -_LOCALE_ALIASES = {"zh": "zh-CN", "zh-hans": "zh-CN", "cn": "zh-CN", "zh-hant": "zh-TW", "tw": "zh-TW", "en-us": "en", "en-gb": "en"} - -@dataclass(frozen=True) -class ObjectLocation: - segments: tuple = () - - @classmethod - def root(cls): - return cls() - - def property(self, name): - if not isinstance(name, str) or not name: - raise ValueError("A canonical KSML property name is required") - return ObjectLocation(self.segments + (("property", name),)) - - def index(self, value): - if value < 0: - raise ValueError("Object location index must not be negative") - return ObjectLocation(self.segments + (("index", value),)) - - def prefixed_by(self, prefix): - return ObjectLocation(prefix.segments + self.segments) - - @builtins.property - def model_path(self): - result = "" - for kind, value in self.segments: - result += f"[{value}]" if kind == "index" else ("." if result else "") + value - return result - - @builtins.property - def native_path(self): - # Python's generated API uses canonical snake_case property names. - return self.model_path - - @builtins.property - def instance_path(self): - def lower_camel(value): - parts = value.split("_") - return parts[0] + "".join(part[:1].upper() + part[1:] for part in parts[1:]) - def escape(value): - return str(value).replace("~", "~0").replace("/", "~1") - return "".join("/" + escape(lower_camel(value) if kind == "property" else value) - for kind, value in self.segments) - - def __str__(self): - return self.native_path - -@dataclass -class CheckResult: - rule_id: str - location: object - input_value: object = None - system_value: object = None - message: str = None - -class CheckException(Exception): - def __init__(self, violations): - self.violations = list(violations) - super().__init__("Check failed: " + "; ".join( - result.message or f"{result.rule_id}:{result.location}" - for result in self.violations)) - -@dataclass(frozen=True) -class ContextEntityRef: - entity: str - id: int - -@dataclass(frozen=True) -class FixEvidence: - entity_type: str - model_path: str - source: str - source_label: str - -class ContextRootError(Exception): - def __init__(self, reason, expected_type, active_root=None): - self.reason, self.expected_type, self.active_root = reason, expected_type, active_root - super().__init__(f"context root {reason}: expected {expected_type}") - -@dataclass(frozen=True, order=True) -class EntityKey: - entity: str - id: object - -class EntityChangeSet: - def __init__(self): self._changes = {} - def set(self, key, field, value): self._changes.setdefault(key, {})[field] = value - def changes(self): return tuple((key, dict(values)) for key, values in self._changes.items()) - def clear_entity(self, key): self._changes.pop(key, None) - def merge_from(self, other): - for key, values in other.changes(): - for field, value in values.items(): self.set(key, field, value) - def rekey(self, old_key, new_key): - values = self._changes.pop(old_key, None) - if values: self._changes.setdefault(new_key, {}).update(values) - -class EntityRoot: - def __init__(self): - self._changes = EntityChangeSet(); self._versions = {}; self._new = set(); self._deleted = set() - def current_change_set(self): return self._changes - def set(self, key, field, value): self._changes.set(key, field, value) - def mark_as_new(self, key): self._new.add(key) - def mark_as_deleted(self, key): self._changes.clear_entity(key); self._deleted.add(key) - def set_original_version(self, key, version): self._versions[key] = version - def original_version(self, key): return self._versions.get(key) - def merge_from(self, other): - if other is self: return - self._changes.merge_from(other._changes); self._versions.update(other._versions) - self._new.update(other._new); self._deleted.update(other._deleted) - def rekey(self, old_key, new_key): - self._changes.rekey(old_key, new_key) - if old_key in self._versions: self._versions[new_key] = self._versions.pop(old_key) - if old_key in self._new: self._new.remove(old_key); self._new.add(new_key) - if old_key in self._deleted: self._deleted.remove(old_key); self._deleted.add(new_key) - def clear_entity(self, key): - self._changes.clear_entity(key); self._new.discard(key); self._deleted.discard(key) - -class SqlLogOperation(str, Enum): - Select = "select" - Insert = "insert" - Update = "update" - Delete = "delete" - -@dataclass(frozen=True) -class SqlLogEntry: - operation: SqlLogOperation - comment: object - purpose: object - audit_reason: object - trace_path: tuple - sql: str - params: tuple - debug_sql: str - elapsed: timedelta - result_count: object = None - affected_rows: object = None - result_summary: str = "" - -class DiagnosticSqlLogSink: - """Value-bearing diagnostic SQL destination; the text sink is installed by default.""" - def write(self, entry): - raise NotImplementedError - -class TextDiagnosticSqlLogSink(DiagnosticSqlLogSink): - def __init__(self, writer=print): self._writer = writer - def write(self, entry): - trace = " -> ".join(f"{key}:{value}" for key, value in entry.trace_path) - comment = "" if entry.comment is None else str(entry.comment) - purpose = "" if entry.purpose is None else str(entry.purpose) - audit_reason = "" if entry.audit_reason is None else str(entry.audit_reason) - self._writer( - f"[TeaQL SQL][{entry.operation.value}][{int(entry.elapsed.total_seconds() * 1000000)}us] " - f"{entry.result_summary} comment={comment} purpose={purpose} " - f"auditReason={audit_reason} tracePath=[{trace}]\n" - f"Parameterized SQL: {entry.sql} params={entry.params!r}\n" - f"Debug SQL: {entry.debug_sql}") - -def _diagnostic_sql_literal(value): - if value is None: return "NULL" - if isinstance(value, bool): return "1" if value else "0" - if isinstance(value, (int, float)): return str(value) - if isinstance(value, (bytes, bytearray)): return "X'" + bytes(value).hex().upper() + "'" - return "'" + str(value).replace("'", "''") + "'" - -def _render_diagnostic_sql(sql, params): - rendered = sql - for value in params: - rendered = rendered.replace("?", _diagnostic_sql_literal(value), 1) - return rendered - -@dataclass(frozen=True) -class RawAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - changes: tuple - -@dataclass(frozen=True) -class SafeAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - fields: tuple - -class UserContext: - """Runtime dependencies and trusted request state initialized by the server.""" - - def __init__(self): - self._resources = {} - self._entity_root = EntityRoot() - self._standard_audit_sink = None - self._app_audit_sink = None - self._audit_policies = {} - self._entity_initializers = {} - self._managed_entities = [] - self._continuous_page_cursors = {} - self._continuous_page_plan = "DISABLED" - self._continuous_page_cursor_id = None - self._id_set_plan = "ID_SET_DISABLED" - self._id_set_count = 0 - self._id_set_count_accuracy = "UNKNOWN" - self._query_sql_log_enabled = True - self._mutation_sql_log_enabled = True - self._sql_logs = [] - self._resources["diagnostic_sql_log_sink"] = TextDiagnosticSqlLogSink() - self._checker_registry = {} - self._checked_mutations = set() - self._graph_save_active = False - self._graph_save_lock = asyncio.Lock() - self._graph_save_owner = contextvars.ContextVar( - f"teaql_graph_save_owner_{id(self)}", default=None) - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - - def begin_fix_evidence(self): - self._resources["fix_evidence_current"] = [] - return self - - def record_fix_evidence(self, entity_type, model_path, source, source_label): - normalized = str(source_label).lower() - if not entity_type or not model_path or source not in ("clock", "context") or not source_label or "authorization" in normalized or "cookie" in normalized or "token=" in normalized: - raise ValueError("Fix evidence must contain only safe framework provenance labels") - self._resources.setdefault("fix_evidence_current", []).append(FixEvidence(entity_type, model_path, source, source_label)) - return self - - def finish_fix_evidence(self): - self._resources["fix_evidence_last"] = tuple(self._resources.get("fix_evidence_current", ())) - self._resources.pop("fix_evidence_current", None) - return self - - def last_fix_evidence(self): - return self._resources.get("fix_evidence_last", ()) - - @classmethod - def new(cls): - return cls() - - def entity_root(self): - return self._entity_root - - def insert_resource(self, resource_type, resource): - self._resources[resource_type] = resource - return self - - def set_locale_code(self, code): - if not isinstance(code, str) or not code.strip(): - raise ValueError(f"Unsupported locale: {code}") - normalized = code.strip().replace("_", "-") - canonical = next((value for value in _SUPPORTED_LOCALES if value.lower() == normalized.lower()), None) - canonical = canonical or _LOCALE_ALIASES.get(normalized.lower()) - if canonical is None: - raise ValueError(f"Unsupported locale: {code}") - self.insert_resource("locale", canonical) - return self - - def set_language_code(self, code): - return self.set_locale_code(code) - - def _translate_check_results(self, results): - locale = self.get_resource("locale") or "en" - messages = _CHECK_MESSAGES.get(locale, _CHECK_MESSAGES["en"]) - for result in results: - key = str(result.rule_id).lower() - if key == "min_str_len": key = "min_length" - if key == "max_str_len": key = "max_length" - template = messages.get(key) or _CHECK_MESSAGES["en"].get(key) or f"checker.{key}" - location = getattr(result.location, "native_path", str(result.location)) - result.message = template.replace("{location}", location) - return results - - async def execute_graph_save(self, work): - if self._graph_save_owner.get() is not None: - return await work() - async with self._graph_save_lock: - provider = self.require_resource("dataService") - begin = getattr(provider, "begin", None) - if not callable(begin): - raise RuntimeError("Configured dataService does not support graph transactions") - transaction = await begin(self) - owner_token = self._graph_save_owner.set(object()) - self._graph_save_active = True - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - from datetime import datetime - self.insert_resource("fix_time", datetime.now()) - self.begin_fix_evidence() - self.insert_resource("dataService", transaction) - try: - result = await work() - except BaseException: - try: - await transaction.rollback(self) - finally: - for action in reversed(self._graph_rollback_actions): - action() - raise - else: - try: - await transaction.commit(self) - except BaseException: - try: - await transaction.rollback(self) - finally: - for action in reversed(self._graph_rollback_actions): - action() - raise - for action in self._graph_commit_actions: - action() - return result - finally: - self.insert_resource("dataService", provider) - self._graph_save_active = False - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - self._resources.pop("fix_time", None) - self.finish_fix_evidence() - self._graph_save_owner.reset(owner_token) - - def after_graph_commit(self, work): - if not self._graph_save_active: - raise RuntimeError("No graph save is active") - self._graph_commit_actions.append(work) - - def after_graph_rollback(self, work): - if not self._graph_save_active: - raise RuntimeError("No graph save is active") - self._graph_rollback_actions.append(work) - - def install(self, module): - """Install a passive metadata manifest; this never changes a database schema.""" - module.apply_to(self) - return self - - async def ensure_schema(self): - """Explicitly reconcile schema and generated bootstrap data.""" - provider = self.require_resource("dataService") - await provider._ensure_schema(self, _SCHEMA_INVOCATION) - - def check_and_fix_mutation(self, mutation): - checker = self._checker_registry.get(getattr(mutation, "entity", None)) - if checker is None: - return - raw_record = getattr(mutation, "payload", getattr(mutation, "values", None)) - if raw_record is None: - return - from datetime import datetime - from teaql.core.value import Value - record = {name: Value.from_any(value) for name, value in raw_record.items()} - owns_fix_time = self._resources.get("fix_time") is None - if owns_fix_time: - self.insert_resource("fix_time", datetime.now()) - self.begin_fix_evidence() - self.insert_resource("fix_operation", "insert" if hasattr(mutation, "payload") else "update") - results = [] - try: - checker.check_and_fix(self, record, None, results) - finally: - if owns_fix_time: - self._resources.pop("fix_time", None) - self.finish_fix_evidence() - self._resources.pop("fix_operation", None) - if results: - self._translate_check_results(results) - raise CheckException(results) - raw_record.clear() - raw_record.update({name: getattr(value, "val", value) for name, value in record.items()}) - - def mark_mutation_checked(self, mutation): - self._checked_mutations.add(id(mutation)) - - def consume_mutation_checked(self, mutation): - key = id(mutation) - if key not in self._checked_mutations: - return False - self._checked_mutations.remove(key) - return True - - def get_resource(self, resource_type): - return self._resources.get(resource_type) - - def require_resource(self, resource_type): - resource = self.get_resource(resource_type) - if resource is None: - raise RuntimeError(f"Required UserContext resource is missing: {resource_type}") - return resource - - def with_active_root(self, root): - if not isinstance(root, ContextEntityRef): - raise TypeError("active root must be ContextEntityRef") - return self.insert_resource("active_root", root) - - def require_active_root(self, expected_type): - root = self.get_resource("active_root") - if not isinstance(root, ContextEntityRef): - raise ContextRootError("missing", expected_type) - if root.entity != expected_type: - raise ContextRootError("type_mismatch", expected_type, root) - return root - - def with_request_policy(self, policy): - self.insert_resource("request_policy", policy) - return self - - def prepare_query(self, query): - policy = self.get_resource("request_policy") - if policy is None: return query - if callable(policy): prepared = policy(query) - elif hasattr(policy, "apply"): prepared = policy.apply(query) - else: raise TypeError("request_policy must be callable or expose apply(query)") - return query if prepared is None else prepared - - def register_entity_initializer(self, entity_name, initializer): - if not isinstance(entity_name, str) or not entity_name.strip() or not callable(initializer): - raise ValueError("entity_name and callable initializer are required") - self._entity_initializers.setdefault(entity_name, []).append(initializer) - return self - - def initialize_entity(self, entity_name, entity): - if not isinstance(entity_name, str) or not entity_name.strip() or entity is None: - raise ValueError("entity_name and entity are required") - for initializer in self._entity_initializers.get("*", ()): - initializer(self, entity) - for initializer in self._entity_initializers.get(entity_name, ()): - initializer(self, entity) - self._managed_entities.append(entity) - return entity - - def managed_entities(self): - return list(self._managed_entities) - - def continuous_page_cursor(self, query_key, offset): - cursor = self._continuous_page_cursors.get((query_key, offset)) - if cursor is not None and cursor["expires_at"] <= __import__("time").time(): - self._continuous_page_cursors.pop((query_key, offset), None) - return None - return cursor - - def put_continuous_page_cursor(self, query_key, offset, cursor): - if len(self._continuous_page_cursors) >= 4096: - oldest = min(self._continuous_page_cursors, - key=lambda key: self._continuous_page_cursors[key]["expires_at"]) - self._continuous_page_cursors.pop(oldest, None) - self._continuous_page_cursors[(query_key, offset)] = cursor - - def observe_continuous_page(self, plan, cursor_id=None): - self._continuous_page_plan = plan - self._continuous_page_cursor_id = cursor_id - - def continuous_page_plan(self): return self._continuous_page_plan - def continuous_page_cursor_id(self): return self._continuous_page_cursor_id - - def id_set_get(self, key): - retained = _ID_SET_STORE.get(key) - if retained is not None and retained["expires_at"] <= time.time(): - _ID_SET_STORE.pop(key, None) - return None - return retained - - def id_set_put(self, key, ids, ttl_seconds): - if len(ids) * 8 > 256 * 1024 * 1024: - raise ValueError("retained ID set exceeds store memory ceiling") - while len(_ID_SET_STORE) >= 64: - oldest = min(_ID_SET_STORE, key=lambda item: _ID_SET_STORE[item]["expires_at"]) - _ID_SET_STORE.pop(oldest, None) - _ID_SET_STORE[key] = {"ids": tuple(ids), "expires_at": time.time() + ttl_seconds} - - def id_set_lock(self, key): - return _ID_SET_LOCKS.setdefault((id(asyncio.get_running_loop()), key), asyncio.Lock()) - - def observe_id_set(self, plan, accuracy="UNKNOWN", count=0): - self._id_set_plan, self._id_set_count_accuracy, self._id_set_count = plan, accuracy, count - - def id_set_plan(self): return self._id_set_plan - def id_set_count(self): return self._id_set_count, self._id_set_count_accuracy - - def _set_sql_log_mode(self, mode): - self._query_sql_log_enabled = mode in ("all", "select") - self._mutation_sql_log_enabled = mode in ("all", "mutation") - self._sql_logs = [] - return self - - def enable_all_sql_log(self): return self._set_sql_log_mode("all") - def enable_select_sql_log(self): return self._set_sql_log_mode("select") - def enable_mutation_sql_log(self): return self._set_sql_log_mode("mutation") - def disable_sql_log(self): return self._set_sql_log_mode("disabled") - def disable_select_sql_log(self): - self._query_sql_log_enabled = False - return self - def disable_mutation_sql_log(self): - self._mutation_sql_log_enabled = False - return self - def clear_sql_logs(self): self._sql_logs = [] - def sql_logs(self): return list(self._sql_logs) - def with_diagnostic_sql_log_sink(self, sink): - self._resources["diagnostic_sql_log_sink"] = sink - return self - def set_diagnostic_sql_log_sink(self, sink): - self._resources["diagnostic_sql_log_sink"] = sink - - def record_sql_evidence(self, operation, sql, params, elapsed_micros, - result_count=None, affected_rows=None, comment=None, - purpose=None, audit_reason=None, trace_path=()): - is_select = operation == SqlLogOperation.Select - if ((is_select and not self._query_sql_log_enabled) - or (not is_select and not self._mutation_sql_log_enabled)): - return - summary = (f"{result_count} rows returned" if result_count is not None - else f"{affected_rows} rows affected") - entry = SqlLogEntry(operation, comment, purpose, audit_reason, tuple(trace_path), - sql, tuple(params), _render_diagnostic_sql(sql, params), - timedelta(microseconds=elapsed_micros), result_count, affected_rows, summary) - self._sql_logs.append(entry) - sink = self._resources.get("diagnostic_sql_log_sink") - if sink is not None: sink.write(entry) - - def initialize_audit(self, standard_sink, app_sink=None): - self._standard_audit_sink = standard_sink - self._app_audit_sink = app_sink - return self - - def configure_audit_policy(self, entity, mask_fields=(), max_length=None): - self._audit_policies[entity] = (frozenset(mask_fields), max_length) - return self - - async def emit_mutation_audit(self, req, result): - command = req.cmd - values = getattr(command, "payload", getattr(command, "values", {})) - kind = "created" if hasattr(command, "payload") else "updated" if hasattr(command, "values") else "deleted" - raw = RawAuditEvent(kind, command.entity, result.get("id"), req.comment, - tuple((name, None, value) for name, value in values.items())) - if self._standard_audit_sink is not None: - emitted = self._standard_audit_sink.on_event(self, raw) - if hasattr(emitted, "__await__"): await emitted - if self._app_audit_sink is not None: - masks, limit = self._audit_policies.get(command.entity, (frozenset(), None)) - fields = [] - for name, _, raw_value in raw.changes: - value = None if raw_value is None else str(raw_value) - masked = name in masks - if value is not None and masked: - value = "*" * len(value) if len(value) < 8 else value[:2] + "*" * (len(value) - 4) + value[-2:] - truncated = value is not None and limit is not None and len(value) > limit - if truncated: value = "*" * limit if limit <= 3 else value[:limit - 3] + "..." - fields.append((name, value, masked, truncated)) - safe = SafeAuditEvent(kind, command.entity, result.get("id"), req.comment, tuple(fields)) - emitted = self._app_audit_sink.on_safe_event(self, safe) - if hasattr(emitted, "__await__"): await emitted - -class RuntimeModule: - """Immutable generated runtime manifest.""" - def __init__(self, entities=(), schemas=None, checkers=None, root_graphs=(), initial_graphs=()): - self.entities = tuple(entities) - self.schemas = dict(schemas or {}) - self.checkers = dict(checkers or {}) - self.root_graphs = tuple(root_graphs) - self.initial_graphs = tuple(initial_graphs) - - def entity(self, entity): - self.entities = (*self.entities, entity) - return self - - def checker(self, entity, checker): - self.checkers[entity] = checker - return self - - def root_graph(self, graph): - self.root_graphs = (*self.root_graphs, graph) - return self - - def initial_graph(self, graph): - self.initial_graphs = (*self.initial_graphs, graph) - return self - - def and_module(self, other): - return RuntimeModule(self.entities + other.entities, - {**self.schemas, **other.schemas}, - {**self.checkers, **other.checkers}, - self.root_graphs + other.root_graphs, - self.initial_graphs + other.initial_graphs) - - def apply_to(self, context): - context.insert_resource("entities", self.entities) - context.insert_resource("entity_schemas", dict(self.schemas)) - context._checker_registry.update(self.checkers) - context.insert_resource("root_graphs", self.root_graphs) - context.insert_resource("initial_graphs", self.initial_graphs) \ No newline at end of file diff --git a/examples/order-management/model.xml b/examples/order-management/model.xml new file mode 100644 index 0000000..3003592 --- /dev/null +++ b/examples/order-management/model.xml @@ -0,0 +1,13 @@ + + + + + + <_value id="1001" name="Pending" code="PENDING" color="#F59E0B" display_order="1" commerce_platform="1"/> + <_value id="1002" name="Confirmed" code="CONFIRMED" color="#10B981" display_order="2" commerce_platform="1"/> + + + + + + diff --git a/examples/order-management/python-app-console/app.py b/examples/order-management/python-app-console/app.py index 27326fb..bfab71f 100644 --- a/examples/order-management/python-app-console/app.py +++ b/examples/order-management/python-app-console/app.py @@ -1,6 +1,7 @@ import asyncio from datetime import date, datetime from decimal import Decimal +import os from pathlib import Path import sys @@ -13,13 +14,14 @@ from models.customer_order import CustomerOrder from models.order_search_preset import OrderSearchPreset from models.order_status import OrderStatus +from runtime_module import GENERATED_RUNTIME_MODULE from teaql.data_service import SQLiteTeaQLClient from teaql.runtime import UserContext class StandardAudit: async def on_event(self, _ctx, event): - print(f"[audit/immutable] {event.kind} {event.entity}#{event.entity_id} reason={event.reason!r}") + print(f"[audit/immutable] {event.kind} {event.entity}#{event.entity_id}") class AppAudit: @@ -33,24 +35,27 @@ async def save(entity, context, reason): async def seed(context): - existing = await (Q.commerce_platforms() - .with_name_is("Northwind Demo") + existing = await (Q.customer_orders() + .with_order_number_is("WEB-2026-001") .comment("Check whether deterministic quick-start data exists") .purpose("Initialize the local order-management example") - .execute_entities_for_list(context)) + .execute_for_list(context)) if existing: print("[seed] deterministic data already exists; no duplicate rows added") return now = datetime(2026, 8, 13, 9, 0, 0) - platform = await save(CommercePlatform(name="Northwind Demo", createTime=now, updateTime=now), context, - "Create quick-start commerce platform") + platform = await (Q.commerce_platforms().with_id_is(1) + .comment("Load generated commerce root").purpose("Seed quick-start data") + .execute_for_one(context)) + assert platform is not None customer = await save(Customer(name="Acme Retail", email="masked-in-quick-start", commercePlatform=platform.id, createTime=now, updateTime=now), context, "Create masked quick-start customer") - pending = await save(OrderStatus(name="Pending", code="PENDING", color="#F97316", - displayOrder=10, commercePlatform=platform.id), context, - "Create quick-start pending status") + pending = await (Q.order_statuses().with_id_is(1001) + .comment("Load generated pending status").purpose("Seed quick-start data") + .execute_for_one(context)) + assert pending is not None await save(CustomerOrder(orderNumber="WEB-2026-001", orderDate=date(2026, 8, 12), totalAmount=Decimal("129.95"), status=pending.id, customer=customer.id, commercePlatform=platform.id, @@ -60,12 +65,12 @@ async def seed(context): async def main(): - database = ROOT / ".local" / "order.db" + database = Path(os.environ.get("TEAQL_ORDER_MANAGEMENT_DB", ROOT / ".local" / "order.db")) if not database.exists(): print(f"[database] {database} was not found; TeaQL will create it") database.parent.mkdir(parents=True, exist_ok=True) client = SQLiteTeaQLClient(str(database)) - context = (UserContext.new() + context = (UserContext.new().install(GENERATED_RUNTIME_MODULE) .insert_resource("dataService", client) .initialize_audit(StandardAudit(), AppAudit()) .configure_audit_policy("Customer", mask_fields=("email",)) @@ -80,21 +85,21 @@ async def main(): .comment("List WEB orders for the terminal quick start") .purpose("Show the operator a deterministic order list") .execute_for_list(context)) - rows = result["data"] + rows = result print(f"[query] matched {len(rows)} order(s)") for row in rows: - print(f" {row['order_number']} {row['order_date']} {row['total_amount']}") + print(f" {row.orderNumber} {row.orderDate} {row.totalAmount}") request_id = "quick-start-pending-orders" preset = await (Q.order_search_presets() .with_request_id_is(request_id) .comment("Check idempotent quick-start preset") .purpose("Persist the operator's reusable search") - .execute_entity_for_one(context)) + .execute_for_one(context)) if preset is None: preset = OrderSearchPreset(name="Pending web orders", filterJson='{"order_number":"WEB-"}', requestId=request_id, ownerUserId="quick-start-user", - commercePlatform=rows[0]["commerce_platform"], + commercePlatform=rows[0].commercePlatform, createTime=datetime.now(), updateTime=datetime.now()) await save(preset, context, "Save idempotent quick-start search preset") print(f"[mutation] saved preset #{preset.id}") diff --git a/examples/order-management/python-lib-core/E.py b/examples/order-management/python-lib-core/E.py index 7e801dd..6e55664 100644 --- a/examples/order-management/python-lib-core/E.py +++ b/examples/order-management/python-lib-core/E.py @@ -19,7 +19,7 @@ def eval(self): raise self._error return self._value - def or_else(self, fallback): + def or_if_null(self, fallback): value = self.eval() return fallback if value is None else value @@ -36,70 +36,339 @@ def eval(self): raise self._error return self._value - def __getattr__(self, method_name): - def access(): - path = f"{self._path}.{method_name}" if self._path else method_name - if self._error is not None: - return ValueExpression(error=self._error) - if self._value is None: - return ValueExpression(None) - field_name = method_name[:-3] if method_name.endswith("_id") else method_name - loaded = getattr(self._value, "_loaded_fields", set()) - if field_name not in loaded and method_name not in loaded: - return ValueExpression(error=TeaQLNotLoadedError(self._root, path, method_name)) - value = getattr(self._value, field_name) - if callable(value): - value = value() - if isinstance(value, list): - return ListExpression(value, self._root, path) - return ValueExpression(value) - return access + def _path_for(self, field): + return f"{self._path}.{field}" if self._path else field + + def _not_loaded(self, field): + path = self._path_for(field) + return TeaQLNotLoadedError(self._root, path, field) + + def _scalar(self, field, relation_id=False): + if self._error is not None: + return ValueExpression(error=self._error) + if self._value is None: + return ValueExpression(None) + if field not in getattr(self._value, "_loaded_fields", set()): + return ValueExpression(error=self._not_loaded(field)) + value = getattr(self._value, field) + if relation_id and value is not None and not isinstance(value, (int, str)): + value = getattr(value, "id", None) + return ValueExpression(value) + + def _relation(self, field, expression_type): + path = self._path_for(field) + if self._error is not None: + return expression_type(None, self._root, path, self._error) + if self._value is None: + return expression_type(None, self._root, path) + if field not in getattr(self._value, "_loaded_fields", set()): + return expression_type(None, self._root, path, self._not_loaded(field)) + value = getattr(self._value, field) + if value is not None and isinstance(value, (int, str)): + return expression_type(None, self._root, path, self._not_loaded(field)) + return expression_type(value, self._root, path) class ListExpression: - def __init__(self, values, root, path): + def __init__(self, values, root, path, item_expression, error=None): self._values = values self._root = root self._path = path + self._item_expression = item_expression + self._error = error def size(self): - return ValueExpression(len(self._values)) + return ValueExpression(error=self._error) if self._error else ValueExpression(len(self._values)) def first(self): return self.get(0) def get(self, index): + path = f"{self._path}.get({index})" + if self._error is not None: + return self._item_expression(None, self._root, path, self._error) value = self._values[index] if 0 <= index < len(self._values) else None - return EntityExpression(value, self._root, f"{self._path}.get({index})") + return self._item_expression(value, self._root, path) + + +class CommercePlatformExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def create_time(self): + return self._scalar("createTime") + def update_time(self): + return self._scalar("updateTime") + def version(self): + return self._scalar("version") + def customer_list(self): + path = self._path_for("customer_list") + if self._error is not None: + return ListExpression([], self._root, path, CustomerExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, CustomerExpression) + if "customer_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, CustomerExpression, self._not_loaded("customer_list")) + return ListExpression(getattr(self._value, "_customer_list"), self._root, path, CustomerExpression) + def order_status_list(self): + path = self._path_for("order_status_list") + if self._error is not None: + return ListExpression([], self._root, path, OrderStatusExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, OrderStatusExpression) + if "order_status_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, OrderStatusExpression, self._not_loaded("order_status_list")) + return ListExpression(getattr(self._value, "_order_status_list"), self._root, path, OrderStatusExpression) + def customer_order_list(self): + path = self._path_for("customer_order_list") + if self._error is not None: + return ListExpression([], self._root, path, CustomerOrderExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, CustomerOrderExpression) + if "customer_order_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, CustomerOrderExpression, self._not_loaded("customer_order_list")) + return ListExpression(getattr(self._value, "_customer_order_list"), self._root, path, CustomerOrderExpression) + def product_list(self): + path = self._path_for("product_list") + if self._error is not None: + return ListExpression([], self._root, path, ProductExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, ProductExpression) + if "product_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, ProductExpression, self._not_loaded("product_list")) + return ListExpression(getattr(self._value, "_product_list"), self._root, path, ProductExpression) + def order_line_list(self): + path = self._path_for("order_line_list") + if self._error is not None: + return ListExpression([], self._root, path, OrderLineExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, OrderLineExpression) + if "order_line_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, OrderLineExpression, self._not_loaded("order_line_list")) + return ListExpression(getattr(self._value, "_order_line_list"), self._root, path, OrderLineExpression) + def order_search_preset_list(self): + path = self._path_for("order_search_preset_list") + if self._error is not None: + return ListExpression([], self._root, path, OrderSearchPresetExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, OrderSearchPresetExpression) + if "order_search_preset_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, OrderSearchPresetExpression, self._not_loaded("order_search_preset_list")) + return ListExpression(getattr(self._value, "_order_search_preset_list"), self._root, path, OrderSearchPresetExpression) + pass + +class CustomerExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def email(self): + return self._scalar("email") + def create_time(self): + return self._scalar("createTime") + def update_time(self): + return self._scalar("updateTime") + def version(self): + return self._scalar("version") + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + def customer_order_list(self): + path = self._path_for("customer_order_list") + if self._error is not None: + return ListExpression([], self._root, path, CustomerOrderExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, CustomerOrderExpression) + if "customer_order_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, CustomerOrderExpression, self._not_loaded("customer_order_list")) + return ListExpression(getattr(self._value, "_customer_order_list"), self._root, path, CustomerOrderExpression) + pass + +class OrderStatusExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def code(self): + return self._scalar("code") + def color(self): + return self._scalar("color") + def display_order(self): + return self._scalar("displayOrder") + def version(self): + return self._scalar("version") + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + def customer_order_list(self): + path = self._path_for("customer_order_list") + if self._error is not None: + return ListExpression([], self._root, path, CustomerOrderExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, CustomerOrderExpression) + if "customer_order_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, CustomerOrderExpression, self._not_loaded("customer_order_list")) + return ListExpression(getattr(self._value, "_customer_order_list"), self._root, path, CustomerOrderExpression) + pass + +class CustomerOrderExpression(EntityExpression): + def id(self): + return self._scalar("id") + def order_number(self): + return self._scalar("orderNumber") + def order_date(self): + return self._scalar("orderDate") + def total_amount(self): + return self._scalar("totalAmount") + def create_time(self): + return self._scalar("createTime") + def update_time(self): + return self._scalar("updateTime") + def version(self): + return self._scalar("version") + def status_id(self): + return self._scalar("status", relation_id=True) + + def status(self): + return self._relation("status", OrderStatusExpression) + def customer_id(self): + return self._scalar("customer", relation_id=True) + + def customer(self): + return self._relation("customer", CustomerExpression) + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + def order_line_list(self): + path = self._path_for("order_line_list") + if self._error is not None: + return ListExpression([], self._root, path, OrderLineExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, OrderLineExpression) + if "order_line_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, OrderLineExpression, self._not_loaded("order_line_list")) + return ListExpression(getattr(self._value, "_order_line_list"), self._root, path, OrderLineExpression) + pass + +class ProductExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def sku(self): + return self._scalar("sku") + def image_url(self): + return self._scalar("imageUrl") + def create_time(self): + return self._scalar("createTime") + def update_time(self): + return self._scalar("updateTime") + def version(self): + return self._scalar("version") + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + def order_line_list(self): + path = self._path_for("order_line_list") + if self._error is not None: + return ListExpression([], self._root, path, OrderLineExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, OrderLineExpression) + if "order_line_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, OrderLineExpression, self._not_loaded("order_line_list")) + return ListExpression(getattr(self._value, "_order_line_list"), self._root, path, OrderLineExpression) + pass + +class OrderLineExpression(EntityExpression): + def id(self): + return self._scalar("id") + def product_name(self): + return self._scalar("productName") + def sku(self): + return self._scalar("sku") + def quantity(self): + return self._scalar("quantity") + def create_time(self): + return self._scalar("createTime") + def version(self): + return self._scalar("version") + def customer_order_id(self): + return self._scalar("customerOrder", relation_id=True) + + def customer_order(self): + return self._relation("customerOrder", CustomerOrderExpression) + def product_id(self): + return self._scalar("product", relation_id=True) + + def product(self): + return self._relation("product", ProductExpression) + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + pass + +class OrderSearchPresetExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def filter_json(self): + return self._scalar("filterJson") + def request_id(self): + return self._scalar("requestId") + def owner_user_id(self): + return self._scalar("ownerUserId") + def create_time(self): + return self._scalar("createTime") + def update_time(self): + return self._scalar("updateTime") + def version(self): + return self._scalar("version") + def commerce_platform_id(self): + return self._scalar("commercePlatform", relation_id=True) + def commerce_platform(self): + return self._relation("commercePlatform", CommercePlatformExpression) + pass class E: @staticmethod def commerce_platform(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "CommercePlatform(id={})".format(entity_id)) + return CommercePlatformExpression(value, "CommercePlatform(id={})".format(entity_id)) @staticmethod def customer(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "Customer(id={})".format(entity_id)) + return CustomerExpression(value, "Customer(id={})".format(entity_id)) @staticmethod def order_status(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "OrderStatus(id={})".format(entity_id)) + return OrderStatusExpression(value, "OrderStatus(id={})".format(entity_id)) @staticmethod def customer_order(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "CustomerOrder(id={})".format(entity_id)) + return CustomerOrderExpression(value, "CustomerOrder(id={})".format(entity_id)) @staticmethod def product(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "Product(id={})".format(entity_id)) + return ProductExpression(value, "Product(id={})".format(entity_id)) @staticmethod def order_line(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "OrderLine(id={})".format(entity_id)) + return OrderLineExpression(value, "OrderLine(id={})".format(entity_id)) @staticmethod def order_search_preset(value): entity_id = getattr(value, "id", None) - return EntityExpression(value, "OrderSearchPreset(id={})".format(entity_id)) + return OrderSearchPresetExpression(value, "OrderSearchPreset(id={})".format(entity_id)) pass \ No newline at end of file diff --git a/examples/order-management/python-lib-core/Q.py b/examples/order-management/python-lib-core/Q.py index 8b41042..ad178bf 100644 --- a/examples/order-management/python-lib-core/Q.py +++ b/examples/order-management/python-lib-core/Q.py @@ -10,29 +10,57 @@ class Q: @staticmethod def commerce_platforms() -> CommercePlatformRequest: - return CommercePlatformRequest() + return CommercePlatformRequest(minimal=False) + + @staticmethod + def commerce_platforms_minimal() -> CommercePlatformRequest: + return CommercePlatformRequest(minimal=True) @staticmethod def customers() -> CustomerRequest: - return CustomerRequest() + return CustomerRequest(minimal=False) + + @staticmethod + def customers_minimal() -> CustomerRequest: + return CustomerRequest(minimal=True) @staticmethod def order_statuses() -> OrderStatusRequest: - return OrderStatusRequest() + return OrderStatusRequest(minimal=False) + + @staticmethod + def order_statuses_minimal() -> OrderStatusRequest: + return OrderStatusRequest(minimal=True) @staticmethod def customer_orders() -> CustomerOrderRequest: - return CustomerOrderRequest() + return CustomerOrderRequest(minimal=False) + + @staticmethod + def customer_orders_minimal() -> CustomerOrderRequest: + return CustomerOrderRequest(minimal=True) @staticmethod def products() -> ProductRequest: - return ProductRequest() + return ProductRequest(minimal=False) + + @staticmethod + def products_minimal() -> ProductRequest: + return ProductRequest(minimal=True) @staticmethod def order_lines() -> OrderLineRequest: - return OrderLineRequest() + return OrderLineRequest(minimal=False) + + @staticmethod + def order_lines_minimal() -> OrderLineRequest: + return OrderLineRequest(minimal=True) @staticmethod def order_search_presets() -> OrderSearchPresetRequest: - return OrderSearchPresetRequest() + return OrderSearchPresetRequest(minimal=False) + + @staticmethod + def order_search_presets_minimal() -> OrderSearchPresetRequest: + return OrderSearchPresetRequest(minimal=True) diff --git a/examples/order-management/python-lib-core/models/commerce_platform.py b/examples/order-management/python-lib-core/models/commerce_platform.py index 9a7e5a7..95b64a0 100644 --- a/examples/order-management/python-lib-core/models/commerce_platform.py +++ b/examples/order-management/python-lib-core/models/commerce_platform.py @@ -1,12 +1,36 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools class CommercePlatform: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "name" in kwargs and "name" not in kwargs: + kwargs["name"] = kwargs.pop("name") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "update_time" in kwargs and "updateTime" not in kwargs: + kwargs["updateTime"] = kwargs.pop("update_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -15,128 +39,399 @@ def __init__(self, **kwargs): self.createTime = kwargs.get("createTime") self.updateTime = kwargs.get("updateTime") self.version = kwargs.get("version") - self._customer_list = [] - self._loaded_fields.add("customer_list") - self._order_status_list = [] - self._loaded_fields.add("order_status_list") - self._customer_order_list = [] - self._loaded_fields.add("customer_order_list") - self._product_list = [] - self._loaded_fields.add("product_list") - self._order_line_list = [] - self._loaded_fields.add("order_line_list") - self._order_search_preset_list = [] - self._loaded_fields.add("order_search_preset_list") + self._customer_list = kwargs.get("customer_list", []) + if "customer_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("customer_list") + self._order_status_list = kwargs.get("order_status_list", []) + if "order_status_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("order_status_list") + self._customer_order_list = kwargs.get("customer_order_list", []) + if "customer_order_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("customer_order_list") + self._product_list = kwargs.get("product_list", []) + if "product_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("product_list") + self._order_line_list = kwargs.get("order_line_list", []) + if "order_line_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("order_line_list") + self._order_search_preset_list = kwargs.get("order_search_preset_list", []) + if "order_search_preset_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("order_search_preset_list") + if self._customer_list: + from models.customer import Customer + self._customer_list = [ + item if isinstance(item, Customer) else Customer(**item) + for item in self._customer_list + ] + if self._order_status_list: + from models.order_status import OrderStatus + self._order_status_list = [ + item if isinstance(item, OrderStatus) else OrderStatus(**item) + for item in self._order_status_list + ] + if self._customer_order_list: + from models.customer_order import CustomerOrder + self._customer_order_list = [ + item if isinstance(item, CustomerOrder) else CustomerOrder(**item) + for item in self._customer_order_list + ] + if self._product_list: + from models.product import Product + self._product_list = [ + item if isinstance(item, Product) else Product(**item) + for item in self._product_list + ] + if self._order_line_list: + from models.order_line import OrderLine + self._order_line_list = [ + item if isinstance(item, OrderLine) else OrderLine(**item) + for item in self._order_line_list + ] + if self._order_search_preset_list: + from models.order_search_preset import OrderSearchPreset + self._order_search_preset_list = [ + item if isinstance(item, OrderSearchPreset) else OrderSearchPreset(**item) + for item in self._order_search_preset_list + ] + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("CommercePlatform", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + for child in self._customer_list: + child._teaql_attach_root(root) + for child in self._order_status_list: + child._teaql_attach_root(root) + for child in self._customer_order_list: + child._teaql_attach_root(root) + for child in self._product_list: + child._teaql_attach_root(root) + for child in self._order_line_list: + child._teaql_attach_root(root) + for child in self._order_search_preset_list: + child._teaql_attach_root(root) + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "name", None) is not None: + if "name" in self._loaded_fields: payload["name"] = Value.Text(self.name) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "updateTime", None) is not None: - payload["update_time"] = Value.Date(self.updateTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "updateTime" in self._loaded_fields: + payload["update_time"] = Value.DateTime(self.updateTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("CommercePlatform", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "CommercePlatform", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "CommercePlatform", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("CommercePlatform", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("CommercePlatform", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "name" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("name"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "updateTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("update_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + for index, child in enumerate(self._customer_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("customer_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + for index, child in enumerate(self._order_status_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("order_status_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + for index, child in enumerate(self._customer_order_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("customer_order_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + for index, child in enumerate(self._product_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("product_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + for index, child in enumerate(self._order_line_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("order_line_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + for index, child in enumerate(self._order_search_preset_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "commercePlatform", self) + child._loaded_fields.add("commercePlatform") + child._entity_root.set(child._teaql_entity_key(), "commerce_platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("order_search_preset_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for CommercePlatform" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + elif "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "update_time" in persisted: + self.updateTime = persisted["update_time"] + self._loaded_fields.add("updateTime") + elif "updateTime" in persisted: + self.updateTime = persisted["updateTime"] + self._loaded_fields.add("updateTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] - cascade_relations.append((self._customer_list, "update_commerce_platform")) - cascade_relations.append((self._order_status_list, "update_commerce_platform")) - cascade_relations.append((self._customer_order_list, "update_commerce_platform")) - cascade_relations.append((self._product_list, "update_commerce_platform")) - cascade_relations.append((self._order_line_list, "update_commerce_platform")) - cascade_relations.append((self._order_search_preset_list, "update_commerce_platform")) + cascade_relations.append(("customer_list", self._customer_list, "update_commerce_platform")) + cascade_relations.append(("order_status_list", self._order_status_list, "update_commerce_platform")) + cascade_relations.append(("customer_order_list", self._customer_order_list, "update_commerce_platform")) + cascade_relations.append(("product_list", self._product_list, "update_commerce_platform")) + cascade_relations.append(("order_line_list", self._order_line_list, "update_commerce_platform")) + cascade_relations.append(("order_search_preset_list", self._order_search_preset_list, "update_commerce_platform")) if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_name(self, value): self.name = value self._loaded_fields.add("name") + self._entity_root.set(self._teaql_entity_key(), "name", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_update_time(self, value): self.updateTime = value self._loaded_fields.add("updateTime") + self._entity_root.set(self._teaql_entity_key(), "update_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def customer_list(self) -> list: + self._loaded_fields.add("customer_list") return self._customer_list def order_status_list(self) -> list: + self._loaded_fields.add("order_status_list") return self._order_status_list def customer_order_list(self) -> list: + self._loaded_fields.add("customer_order_list") return self._customer_order_list def product_list(self) -> list: + self._loaded_fields.add("product_list") return self._product_list def order_line_list(self) -> list: + self._loaded_fields.add("order_line_list") return self._order_line_list def order_search_preset_list(self) -> list: + self._loaded_fields.add("order_search_preset_list") return self._order_search_preset_list \ No newline at end of file diff --git a/examples/order-management/python-lib-core/models/customer.py b/examples/order-management/python-lib-core/models/customer.py index 1889aa0..5029f09 100644 --- a/examples/order-management/python-lib-core/models/customer.py +++ b/examples/order-management/python-lib-core/models/customer.py @@ -1,12 +1,41 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.commerce_platform import CommercePlatform class Customer: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "name" in kwargs and "name" not in kwargs: + kwargs["name"] = kwargs.pop("name") + if "email" in kwargs and "email" not in kwargs: + kwargs["email"] = kwargs.pop("email") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "update_time" in kwargs and "updateTime" not in kwargs: + kwargs["updateTime"] = kwargs.pop("update_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -17,112 +46,283 @@ def __init__(self, **kwargs): self.createTime = kwargs.get("createTime") self.updateTime = kwargs.get("updateTime") self.version = kwargs.get("version") - self._customer_order_list = [] - self._loaded_fields.add("customer_order_list") + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._customer_order_list = kwargs.get("customer_order_list", []) + if "customer_order_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("customer_order_list") + if self._customer_order_list: + from models.customer_order import CustomerOrder + self._customer_order_list = [ + item if isinstance(item, CustomerOrder) else CustomerOrder(**item) + for item in self._customer_order_list + ] + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("Customer", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + for child in self._customer_order_list: + child._teaql_attach_root(root) + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "name", None) is not None: + if "name" in self._loaded_fields: payload["name"] = Value.Text(self.name) - if getattr(self, "email", None) is not None: + if "email" in self._loaded_fields: payload["email"] = Value.Text(self.email) - if getattr(self, "commercePlatform", None) is not None: + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "updateTime", None) is not None: - payload["update_time"] = Value.Date(self.updateTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "updateTime" in self._loaded_fields: + payload["update_time"] = Value.DateTime(self.updateTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("Customer", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "Customer", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "Customer", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("Customer", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("Customer", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "name" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("name"), message="Mutation requires a fully loaded entity")]) + if "email" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("email"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "updateTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("update_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + for index, child in enumerate(self._customer_order_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "customer", self) + child._loaded_fields.add("customer") + child._entity_root.set(child._teaql_entity_key(), "customer", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("customer_order_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for Customer" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + elif "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + if "email" in persisted: + self.email = persisted["email"] + self._loaded_fields.add("email") + elif "email" in persisted: + self.email = persisted["email"] + self._loaded_fields.add("email") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "update_time" in persisted: + self.updateTime = persisted["update_time"] + self._loaded_fields.add("updateTime") + elif "updateTime" in persisted: + self.updateTime = persisted["updateTime"] + self._loaded_fields.add("updateTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] - cascade_relations.append((self._customer_order_list, "update_customer")) + cascade_relations.append(("customer_order_list", self._customer_order_list, "update_customer")) if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_name(self, value): self.name = value self._loaded_fields.add("name") + self._entity_root.set(self._teaql_entity_key(), "name", Value.from_any(value)) return self def update_email(self, value): self.email = value self._loaded_fields.add("email") + self._entity_root.set(self._teaql_entity_key(), "email", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_update_time(self, value): self.updateTime = value self._loaded_fields.add("updateTime") + self._entity_root.set(self._teaql_entity_key(), "update_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self def customer_order_list(self) -> list: + self._loaded_fields.add("customer_order_list") return self._customer_order_list \ No newline at end of file diff --git a/examples/order-management/python-lib-core/models/customer_order.py b/examples/order-management/python-lib-core/models/customer_order.py index bf89bdd..5292dbc 100644 --- a/examples/order-management/python-lib-core/models/customer_order.py +++ b/examples/order-management/python-lib-core/models/customer_order.py @@ -1,12 +1,49 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.order_status import OrderStatus +from models.customer import Customer +from models.commerce_platform import CommercePlatform class CustomerOrder: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "order_number" in kwargs and "orderNumber" not in kwargs: + kwargs["orderNumber"] = kwargs.pop("order_number") + if "order_date" in kwargs and "orderDate" not in kwargs: + kwargs["orderDate"] = kwargs.pop("order_date") + if "total_amount" in kwargs and "totalAmount" not in kwargs: + kwargs["totalAmount"] = kwargs.pop("total_amount") + if "status" in kwargs and "status" not in kwargs: + kwargs["status"] = kwargs.pop("status") + if "customer" in kwargs and "customer" not in kwargs: + kwargs["customer"] = kwargs.pop("customer") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "update_time" in kwargs and "updateTime" not in kwargs: + kwargs["updateTime"] = kwargs.pop("update_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -20,151 +57,345 @@ def __init__(self, **kwargs): self.createTime = kwargs.get("createTime") self.updateTime = kwargs.get("updateTime") self.version = kwargs.get("version") - self._order_line_list = [] - self._loaded_fields.add("order_line_list") + if isinstance(self.status, dict): + self.status = OrderStatus(**self.status) + if isinstance(self.customer, dict): + self.customer = Customer(**self.customer) + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._order_line_list = kwargs.get("order_line_list", []) + if "order_line_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("order_line_list") + if self._order_line_list: + from models.order_line import OrderLine + self._order_line_list = [ + item if isinstance(item, OrderLine) else OrderLine(**item) + for item in self._order_line_list + ] + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("CustomerOrder", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + for child in self._order_line_list: + child._teaql_attach_root(root) + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "orderNumber", None) is not None: + if "orderNumber" in self._loaded_fields: payload["order_number"] = Value.Text(self.orderNumber) - if getattr(self, "orderDate", None) is not None: + if "orderDate" in self._loaded_fields: payload["order_date"] = Value.Date(self.orderDate) - if getattr(self, "totalAmount", None) is not None: - payload["total_amount"] = Value.Object(self.totalAmount) - if getattr(self, "status", None) is not None: + if "totalAmount" in self._loaded_fields: + payload["total_amount"] = Value.Decimal(self.totalAmount) + if "status" in self._loaded_fields: payload["status"] = Value.Object(self.status) - if getattr(self, "customer", None) is not None: + if "customer" in self._loaded_fields: payload["customer"] = Value.Object(self.customer) - if getattr(self, "commercePlatform", None) is not None: + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "updateTime", None) is not None: - payload["update_time"] = Value.Date(self.updateTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "updateTime" in self._loaded_fields: + payload["update_time"] = Value.DateTime(self.updateTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("CustomerOrder", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "CustomerOrder", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "CustomerOrder", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("CustomerOrder", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("CustomerOrder", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "orderNumber" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("order_number"), message="Mutation requires a fully loaded entity")]) + if "orderDate" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("order_date"), message="Mutation requires a fully loaded entity")]) + if "totalAmount" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("total_amount"), message="Mutation requires a fully loaded entity")]) + if "status" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("status"), message="Mutation requires a fully loaded entity")]) + if "customer" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("customer"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "updateTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("update_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + for index, child in enumerate(self._order_line_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "customerOrder", self) + child._loaded_fields.add("customerOrder") + child._entity_root.set(child._teaql_entity_key(), "customer_order", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("order_line_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for CustomerOrder" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "order_number" in persisted: + self.orderNumber = persisted["order_number"] + self._loaded_fields.add("orderNumber") + elif "orderNumber" in persisted: + self.orderNumber = persisted["orderNumber"] + self._loaded_fields.add("orderNumber") + if "order_date" in persisted: + self.orderDate = persisted["order_date"] + self._loaded_fields.add("orderDate") + elif "orderDate" in persisted: + self.orderDate = persisted["orderDate"] + self._loaded_fields.add("orderDate") + if "total_amount" in persisted: + self.totalAmount = persisted["total_amount"] + self._loaded_fields.add("totalAmount") + elif "totalAmount" in persisted: + self.totalAmount = persisted["totalAmount"] + self._loaded_fields.add("totalAmount") + if "status" in persisted: + self.status = persisted["status"] + self._loaded_fields.add("status") + elif "status" in persisted: + self.status = persisted["status"] + self._loaded_fields.add("status") + if "customer" in persisted: + self.customer = persisted["customer"] + self._loaded_fields.add("customer") + elif "customer" in persisted: + self.customer = persisted["customer"] + self._loaded_fields.add("customer") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "update_time" in persisted: + self.updateTime = persisted["update_time"] + self._loaded_fields.add("updateTime") + elif "updateTime" in persisted: + self.updateTime = persisted["updateTime"] + self._loaded_fields.add("updateTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] - cascade_relations.append((self._order_line_list, "update_customer_order")) + cascade_relations.append(("order_line_list", self._order_line_list, "update_customer_order")) if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_order_number(self, value): self.orderNumber = value self._loaded_fields.add("orderNumber") + self._entity_root.set(self._teaql_entity_key(), "order_number", Value.from_any(value)) return self def update_order_date(self, value): self.orderDate = value self._loaded_fields.add("orderDate") + self._entity_root.set(self._teaql_entity_key(), "order_date", Value.from_any(value)) return self def update_total_amount(self, value): self.totalAmount = value self._loaded_fields.add("totalAmount") + self._entity_root.set(self._teaql_entity_key(), "total_amount", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_update_time(self, value): self.updateTime = value self._loaded_fields.add("updateTime") + self._entity_root.set(self._teaql_entity_key(), "update_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_status(self, value): self.status = getattr(value, "id", value) if value else None self._loaded_fields.add("status") + self._entity_root.set(self._teaql_entity_key(), "status", Value.from_any(self.status)) return self def update_status_to_pending(self): self.status = 1001 self._loaded_fields.add("status") return self - def update_status_to_processing(self): + def update_status_to_confirmed(self): self.status = 1002 self._loaded_fields.add("status") return self - def update_status_to_shipped(self): - self.status = 1003 - self._loaded_fields.add("status") - return self - def update_status_to_completed(self): - self.status = 1004 - self._loaded_fields.add("status") - return self def update_customer(self, value): self.customer = getattr(value, "id", value) if value else None self._loaded_fields.add("customer") + self._entity_root.set(self._teaql_entity_key(), "customer", Value.from_any(self.customer)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self def order_line_list(self) -> list: + self._loaded_fields.add("order_line_list") return self._order_line_list \ No newline at end of file diff --git a/examples/order-management/python-lib-core/models/order_line.py b/examples/order-management/python-lib-core/models/order_line.py index a6ecb18..fe06410 100644 --- a/examples/order-management/python-lib-core/models/order_line.py +++ b/examples/order-management/python-lib-core/models/order_line.py @@ -1,12 +1,47 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.customer_order import CustomerOrder +from models.product import Product +from models.commerce_platform import CommercePlatform class OrderLine: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "customer_order" in kwargs and "customerOrder" not in kwargs: + kwargs["customerOrder"] = kwargs.pop("customer_order") + if "product" in kwargs and "product" not in kwargs: + kwargs["product"] = kwargs.pop("product") + if "product_name" in kwargs and "productName" not in kwargs: + kwargs["productName"] = kwargs.pop("product_name") + if "sku" in kwargs and "sku" not in kwargs: + kwargs["sku"] = kwargs.pop("sku") + if "quantity" in kwargs and "quantity" not in kwargs: + kwargs["quantity"] = kwargs.pop("quantity") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -19,123 +54,292 @@ def __init__(self, **kwargs): self.commercePlatform = kwargs.get("commercePlatform") self.createTime = kwargs.get("createTime") self.version = kwargs.get("version") + if isinstance(self.customerOrder, dict): + self.customerOrder = CustomerOrder(**self.customerOrder) + if isinstance(self.product, dict): + self.product = Product(**self.product) + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("OrderLine", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "customerOrder", None) is not None: + if "customerOrder" in self._loaded_fields: payload["customer_order"] = Value.Object(self.customerOrder) - if getattr(self, "product", None) is not None: + if "product" in self._loaded_fields: payload["product"] = Value.Object(self.product) - if getattr(self, "productName", None) is not None: + if "productName" in self._loaded_fields: payload["product_name"] = Value.Text(self.productName) - if getattr(self, "sku", None) is not None: + if "sku" in self._loaded_fields: payload["sku"] = Value.Text(self.sku) - if getattr(self, "quantity", None) is not None: - payload["quantity"] = Value.Object(self.quantity) - if getattr(self, "commercePlatform", None) is not None: + if "quantity" in self._loaded_fields: + payload["quantity"] = Value.I64(self.quantity) + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("OrderLine", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "OrderLine", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "OrderLine", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("OrderLine", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("OrderLine", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "customerOrder" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("customer_order"), message="Mutation requires a fully loaded entity")]) + if "product" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("product"), message="Mutation requires a fully loaded entity")]) + if "productName" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("product_name"), message="Mutation requires a fully loaded entity")]) + if "sku" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("sku"), message="Mutation requires a fully loaded entity")]) + if "quantity" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("quantity"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for OrderLine" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "customer_order" in persisted: + self.customerOrder = persisted["customer_order"] + self._loaded_fields.add("customerOrder") + elif "customerOrder" in persisted: + self.customerOrder = persisted["customerOrder"] + self._loaded_fields.add("customerOrder") + if "product" in persisted: + self.product = persisted["product"] + self._loaded_fields.add("product") + elif "product" in persisted: + self.product = persisted["product"] + self._loaded_fields.add("product") + if "product_name" in persisted: + self.productName = persisted["product_name"] + self._loaded_fields.add("productName") + elif "productName" in persisted: + self.productName = persisted["productName"] + self._loaded_fields.add("productName") + if "sku" in persisted: + self.sku = persisted["sku"] + self._loaded_fields.add("sku") + elif "sku" in persisted: + self.sku = persisted["sku"] + self._loaded_fields.add("sku") + if "quantity" in persisted: + self.quantity = persisted["quantity"] + self._loaded_fields.add("quantity") + elif "quantity" in persisted: + self.quantity = persisted["quantity"] + self._loaded_fields.add("quantity") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_product_name(self, value): self.productName = value self._loaded_fields.add("productName") + self._entity_root.set(self._teaql_entity_key(), "product_name", Value.from_any(value)) return self def update_sku(self, value): self.sku = value self._loaded_fields.add("sku") + self._entity_root.set(self._teaql_entity_key(), "sku", Value.from_any(value)) return self def update_quantity(self, value): self.quantity = value self._loaded_fields.add("quantity") + self._entity_root.set(self._teaql_entity_key(), "quantity", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_customer_order(self, value): self.customerOrder = getattr(value, "id", value) if value else None self._loaded_fields.add("customerOrder") + self._entity_root.set(self._teaql_entity_key(), "customer_order", Value.from_any(self.customerOrder)) return self def update_product(self, value): self.product = getattr(value, "id", value) if value else None self._loaded_fields.add("product") + self._entity_root.set(self._teaql_entity_key(), "product", Value.from_any(self.product)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self diff --git a/examples/order-management/python-lib-core/models/order_search_preset.py b/examples/order-management/python-lib-core/models/order_search_preset.py index 00b0971..4e8f370 100644 --- a/examples/order-management/python-lib-core/models/order_search_preset.py +++ b/examples/order-management/python-lib-core/models/order_search_preset.py @@ -1,12 +1,45 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.commerce_platform import CommercePlatform class OrderSearchPreset: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "name" in kwargs and "name" not in kwargs: + kwargs["name"] = kwargs.pop("name") + if "filter_json" in kwargs and "filterJson" not in kwargs: + kwargs["filterJson"] = kwargs.pop("filter_json") + if "request_id" in kwargs and "requestId" not in kwargs: + kwargs["requestId"] = kwargs.pop("request_id") + if "owner_user_id" in kwargs and "ownerUserId" not in kwargs: + kwargs["ownerUserId"] = kwargs.pop("owner_user_id") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "update_time" in kwargs and "updateTime" not in kwargs: + kwargs["updateTime"] = kwargs.pop("update_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -19,121 +52,286 @@ def __init__(self, **kwargs): self.createTime = kwargs.get("createTime") self.updateTime = kwargs.get("updateTime") self.version = kwargs.get("version") + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("OrderSearchPreset", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "name", None) is not None: + if "name" in self._loaded_fields: payload["name"] = Value.Text(self.name) - if getattr(self, "filterJson", None) is not None: + if "filterJson" in self._loaded_fields: payload["filter_json"] = Value.Text(self.filterJson) - if getattr(self, "requestId", None) is not None: + if "requestId" in self._loaded_fields: payload["request_id"] = Value.Text(self.requestId) - if getattr(self, "ownerUserId", None) is not None: + if "ownerUserId" in self._loaded_fields: payload["owner_user_id"] = Value.Text(self.ownerUserId) - if getattr(self, "commercePlatform", None) is not None: + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "updateTime", None) is not None: - payload["update_time"] = Value.Date(self.updateTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "updateTime" in self._loaded_fields: + payload["update_time"] = Value.DateTime(self.updateTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("OrderSearchPreset", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "OrderSearchPreset", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "OrderSearchPreset", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("OrderSearchPreset", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("OrderSearchPreset", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "name" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("name"), message="Mutation requires a fully loaded entity")]) + if "filterJson" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("filter_json"), message="Mutation requires a fully loaded entity")]) + if "requestId" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("request_id"), message="Mutation requires a fully loaded entity")]) + if "ownerUserId" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("owner_user_id"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "updateTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("update_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for OrderSearchPreset" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + elif "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + if "filter_json" in persisted: + self.filterJson = persisted["filter_json"] + self._loaded_fields.add("filterJson") + elif "filterJson" in persisted: + self.filterJson = persisted["filterJson"] + self._loaded_fields.add("filterJson") + if "request_id" in persisted: + self.requestId = persisted["request_id"] + self._loaded_fields.add("requestId") + elif "requestId" in persisted: + self.requestId = persisted["requestId"] + self._loaded_fields.add("requestId") + if "owner_user_id" in persisted: + self.ownerUserId = persisted["owner_user_id"] + self._loaded_fields.add("ownerUserId") + elif "ownerUserId" in persisted: + self.ownerUserId = persisted["ownerUserId"] + self._loaded_fields.add("ownerUserId") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "update_time" in persisted: + self.updateTime = persisted["update_time"] + self._loaded_fields.add("updateTime") + elif "updateTime" in persisted: + self.updateTime = persisted["updateTime"] + self._loaded_fields.add("updateTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_name(self, value): self.name = value self._loaded_fields.add("name") + self._entity_root.set(self._teaql_entity_key(), "name", Value.from_any(value)) return self def update_filter_json(self, value): self.filterJson = value self._loaded_fields.add("filterJson") + self._entity_root.set(self._teaql_entity_key(), "filter_json", Value.from_any(value)) return self def update_request_id(self, value): self.requestId = value self._loaded_fields.add("requestId") + self._entity_root.set(self._teaql_entity_key(), "request_id", Value.from_any(value)) return self def update_owner_user_id(self, value): self.ownerUserId = value self._loaded_fields.add("ownerUserId") + self._entity_root.set(self._teaql_entity_key(), "owner_user_id", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_update_time(self, value): self.updateTime = value self._loaded_fields.add("updateTime") + self._entity_root.set(self._teaql_entity_key(), "update_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self diff --git a/examples/order-management/python-lib-core/models/order_status.py b/examples/order-management/python-lib-core/models/order_status.py index df72e63..d71a584 100644 --- a/examples/order-management/python-lib-core/models/order_status.py +++ b/examples/order-management/python-lib-core/models/order_status.py @@ -1,12 +1,41 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.commerce_platform import CommercePlatform class OrderStatus: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "name" in kwargs and "name" not in kwargs: + kwargs["name"] = kwargs.pop("name") + if "code" in kwargs and "code" not in kwargs: + kwargs["code"] = kwargs.pop("code") + if "color" in kwargs and "color" not in kwargs: + kwargs["color"] = kwargs.pop("color") + if "display_order" in kwargs and "displayOrder" not in kwargs: + kwargs["displayOrder"] = kwargs.pop("display_order") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -17,112 +46,283 @@ def __init__(self, **kwargs): self.displayOrder = kwargs.get("displayOrder") self.commercePlatform = kwargs.get("commercePlatform") self.version = kwargs.get("version") - self._customer_order_list = [] - self._loaded_fields.add("customer_order_list") + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._customer_order_list = kwargs.get("customer_order_list", []) + if "customer_order_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("customer_order_list") + if self._customer_order_list: + from models.customer_order import CustomerOrder + self._customer_order_list = [ + item if isinstance(item, CustomerOrder) else CustomerOrder(**item) + for item in self._customer_order_list + ] + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("OrderStatus", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + for child in self._customer_order_list: + child._teaql_attach_root(root) + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "name", None) is not None: + if "name" in self._loaded_fields: payload["name"] = Value.Text(self.name) - if getattr(self, "code", None) is not None: + if "code" in self._loaded_fields: payload["code"] = Value.Text(self.code) - if getattr(self, "color", None) is not None: + if "color" in self._loaded_fields: payload["color"] = Value.Text(self.color) - if getattr(self, "displayOrder", None) is not None: - payload["display_order"] = Value.Object(self.displayOrder) - if getattr(self, "commercePlatform", None) is not None: + if "displayOrder" in self._loaded_fields: + payload["display_order"] = Value.Decimal(self.displayOrder) + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "version", None) is not None: + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("OrderStatus", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "OrderStatus", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "OrderStatus", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("OrderStatus", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("OrderStatus", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "name" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("name"), message="Mutation requires a fully loaded entity")]) + if "code" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("code"), message="Mutation requires a fully loaded entity")]) + if "color" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("color"), message="Mutation requires a fully loaded entity")]) + if "displayOrder" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("display_order"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + for index, child in enumerate(self._customer_order_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "status", self) + child._loaded_fields.add("status") + child._entity_root.set(child._teaql_entity_key(), "status", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("customer_order_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for OrderStatus" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + elif "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + if "code" in persisted: + self.code = persisted["code"] + self._loaded_fields.add("code") + elif "code" in persisted: + self.code = persisted["code"] + self._loaded_fields.add("code") + if "color" in persisted: + self.color = persisted["color"] + self._loaded_fields.add("color") + elif "color" in persisted: + self.color = persisted["color"] + self._loaded_fields.add("color") + if "display_order" in persisted: + self.displayOrder = persisted["display_order"] + self._loaded_fields.add("displayOrder") + elif "displayOrder" in persisted: + self.displayOrder = persisted["displayOrder"] + self._loaded_fields.add("displayOrder") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] - cascade_relations.append((self._customer_order_list, "update_status")) + cascade_relations.append(("customer_order_list", self._customer_order_list, "update_status")) if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_name(self, value): self.name = value self._loaded_fields.add("name") + self._entity_root.set(self._teaql_entity_key(), "name", Value.from_any(value)) return self def update_code(self, value): self.code = value self._loaded_fields.add("code") + self._entity_root.set(self._teaql_entity_key(), "code", Value.from_any(value)) return self def update_color(self, value): self.color = value self._loaded_fields.add("color") + self._entity_root.set(self._teaql_entity_key(), "color", Value.from_any(value)) return self def update_display_order(self, value): self.displayOrder = value self._loaded_fields.add("displayOrder") + self._entity_root.set(self._teaql_entity_key(), "display_order", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self def customer_order_list(self) -> list: + self._loaded_fields.add("customer_order_list") return self._customer_order_list \ No newline at end of file diff --git a/examples/order-management/python-lib-core/models/product.py b/examples/order-management/python-lib-core/models/product.py index b8859da..d82518a 100644 --- a/examples/order-management/python-lib-core/models/product.py +++ b/examples/order-management/python-lib-core/models/product.py @@ -1,12 +1,43 @@ from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest from teaql.core.value import Value +from teaql.runtime import CheckException, CheckResult, EntityKey, EntityRoot, ObjectLocation +import itertools +from models.commerce_platform import CommercePlatform class Product: + _teaql_temporary_ids = itertools.count(1) @classmethod def refer(cls, entity_id): return cls(id=entity_id) + @classmethod + def _teaql_new_with_fixed_id(cls, entity_id): + """Generated bootstrap capability; application code must not call it.""" + return cls(id=entity_id)._teaql_force_create() + + def _teaql_force_create(self): + self._action = "Create" + self._entity_root.mark_as_new(self._teaql_entity_key()) + return self + def __init__(self, **kwargs): + self._entity_root = kwargs.pop("_entity_root", None) or EntityRoot() + if "id" in kwargs and "id" not in kwargs: + kwargs["id"] = kwargs.pop("id") + if "name" in kwargs and "name" not in kwargs: + kwargs["name"] = kwargs.pop("name") + if "sku" in kwargs and "sku" not in kwargs: + kwargs["sku"] = kwargs.pop("sku") + if "image_url" in kwargs and "imageUrl" not in kwargs: + kwargs["imageUrl"] = kwargs.pop("image_url") + if "commerce_platform" in kwargs and "commercePlatform" not in kwargs: + kwargs["commercePlatform"] = kwargs.pop("commerce_platform") + if "create_time" in kwargs and "createTime" not in kwargs: + kwargs["createTime"] = kwargs.pop("create_time") + if "update_time" in kwargs and "updateTime" not in kwargs: + kwargs["updateTime"] = kwargs.pop("update_time") + if "version" in kwargs and "version" not in kwargs: + kwargs["version"] = kwargs.pop("version") self._action = "Update" if kwargs.get("id") else "Create" self._comment = None self._loaded_fields = set(kwargs.keys()) @@ -18,119 +49,299 @@ def __init__(self, **kwargs): self.createTime = kwargs.get("createTime") self.updateTime = kwargs.get("updateTime") self.version = kwargs.get("version") - self._order_line_list = [] - self._loaded_fields.add("order_line_list") + if isinstance(self.commercePlatform, dict): + self.commercePlatform = CommercePlatform(**self.commercePlatform) + self._order_line_list = kwargs.get("order_line_list", []) + if "order_line_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("order_line_list") + if self._order_line_list: + from models.order_line import OrderLine + self._order_line_list = [ + item if isinstance(item, OrderLine) else OrderLine(**item) + for item in self._order_line_list + ] + self._ledger_id = getattr(self, "id", None) + if self._ledger_id is None: + self._ledger_id = -next(self._teaql_temporary_ids) + key = self._teaql_entity_key() + if self._action == "Create": + self._entity_root.mark_as_new(key) + elif getattr(self, "version", None) is not None: + self._entity_root.set_original_version(key, int(self.version)) + + def _teaql_entity_key(self): + return EntityKey("Product", self._ledger_id) + + def _teaql_attach_root(self, root): + if self._entity_root is not root: + root.merge_from(self._entity_root) + self._entity_root = root + for child in self._order_line_list: + child._teaql_attach_root(root) + return self def mark_for_deletion(self): self._action = "Delete" + self._entity_root.mark_as_deleted(self._teaql_entity_key()) return self def audit_as(self, comment: str): + if not isinstance(comment, str) or not comment.strip(): + raise ValueError("Security audit failure: audit_as() requires a non-empty reason") self._comment = comment return self async def save(self, context): - if not self._comment: - raise Exception("Security audit failure: audit_as() must be called before save()") + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) + + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) + def _teaql_build_command(self): payload = {} - if getattr(self, "id", None) is not None: + if "id" in self._loaded_fields: payload["id"] = Value.I64(self.id) - if getattr(self, "name", None) is not None: + if "name" in self._loaded_fields: payload["name"] = Value.Text(self.name) - if getattr(self, "sku", None) is not None: + if "sku" in self._loaded_fields: payload["sku"] = Value.Text(self.sku) - if getattr(self, "imageUrl", None) is not None: + if "imageUrl" in self._loaded_fields: payload["image_url"] = Value.Text(self.imageUrl) - if getattr(self, "commercePlatform", None) is not None: + if "commercePlatform" in self._loaded_fields: payload["commerce_platform"] = Value.Object(self.commercePlatform) - if getattr(self, "createTime", None) is not None: - payload["create_time"] = Value.Date(self.createTime) - if getattr(self, "updateTime", None) is not None: - payload["update_time"] = Value.Date(self.updateTime) - if getattr(self, "version", None) is not None: + if "createTime" in self._loaded_fields: + payload["create_time"] = Value.DateTime(self.createTime) + if "updateTime" in self._loaded_fields: + payload["update_time"] = Value.DateTime(self.updateTime) + if "version" in self._loaded_fields: payload["version"] = Value.I64(self.version) - action = self._action + if action == "Update": + ledger = dict(self._entity_root.current_change_set().changes()).get(self._teaql_entity_key(), {}) + payload = {field: value for field, value in ledger.items() if field not in ("id", "version")} if action == "Create": cmd = InsertCommand("Product", payload) - elif self._action == "Update": - cmd = UpdateCommand( - "Product", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) - for k, v in payload.items(): - if k not in ("id", "version"): - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand( - "Product", - Value.from_any(getattr(self, "id", None)), - getattr(self, "version", None), - ) + elif action == "Update": + cmd = UpdateCommand("Product", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + for key, value in payload.items(): + if key not in ("id", "version"): cmd.value(key, value) + else: + cmd = DeleteCommand("Product", Value.from_any(getattr(self, "id", None)), getattr(self, "version", None)) + return action, cmd + + def _teaql_preflight_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + if self._action == "Update": + if "id" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("id"), message="Mutation requires a fully loaded entity")]) + if "name" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("name"), message="Mutation requires a fully loaded entity")]) + if "sku" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("sku"), message="Mutation requires a fully loaded entity")]) + if "imageUrl" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("image_url"), message="Mutation requires a fully loaded entity")]) + if "commercePlatform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("commerce_platform"), message="Mutation requires a fully loaded entity")]) + if "createTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("create_time"), message="Mutation requires a fully loaded entity")]) + if "updateTime" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("update_time"), message="Mutation requires a fully loaded entity")]) + if "version" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("version"), message="Mutation requires a fully loaded entity")]) + _action, cmd = self._teaql_build_command() + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + for index, child in enumerate(self._order_line_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "product", self) + child._loaded_fields.add("product") + child._entity_root.set(child._teaql_entity_key(), "product", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("order_line_list").index(index) + raise CheckException([ + CheckResult(v.rule_id, v.location.prefixed_by(prefix), v.input_value, v.system_value, v.message) + for v in error.violations + ]) from error + + async def _teaql_save_within_graph(self, context): + if not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: audit_as() must be called before save()") + + self._teaql_attach_root(self._entity_root) + action, cmd = self._teaql_build_command() req = MutationRequest(cmd) if self._comment: req.comment = self._comment + try: + context.check_and_fix_mutation(cmd) + finally: + for field, value in getattr(cmd, "values", {}).items(): + if field not in ("id", "version"): + self._entity_root.set(self._teaql_entity_key(), field, value) + context.mark_mutation_checked(cmd) service = context.require_resource("dataService") result = await service.mutate(context, req) - if action == "Create": - self.id = result["id"] - self.version = result.get("version") + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for Product" + ) + rollback_payload = {field: getattr(self, field, None) for field in self._loaded_fields | {"id", "version"}} + rollback_ledger_id = self._ledger_id + rollback_action = self._action + rollback_loaded_fields = set(self._loaded_fields) + old_key = self._teaql_entity_key() + if "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + elif "id" in persisted: + self.id = persisted["id"] + self._loaded_fields.add("id") + if "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + elif "name" in persisted: + self.name = persisted["name"] + self._loaded_fields.add("name") + if "sku" in persisted: + self.sku = persisted["sku"] + self._loaded_fields.add("sku") + elif "sku" in persisted: + self.sku = persisted["sku"] + self._loaded_fields.add("sku") + if "image_url" in persisted: + self.imageUrl = persisted["image_url"] + self._loaded_fields.add("imageUrl") + elif "imageUrl" in persisted: + self.imageUrl = persisted["imageUrl"] + self._loaded_fields.add("imageUrl") + if "commerce_platform" in persisted: + self.commercePlatform = persisted["commerce_platform"] + self._loaded_fields.add("commercePlatform") + elif "commercePlatform" in persisted: + self.commercePlatform = persisted["commercePlatform"] + self._loaded_fields.add("commercePlatform") + if "create_time" in persisted: + self.createTime = persisted["create_time"] + self._loaded_fields.add("createTime") + elif "createTime" in persisted: + self.createTime = persisted["createTime"] + self._loaded_fields.add("createTime") + if "update_time" in persisted: + self.updateTime = persisted["update_time"] + self._loaded_fields.add("updateTime") + elif "updateTime" in persisted: + self.updateTime = persisted["updateTime"] + self._loaded_fields.add("updateTime") + if "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + elif "version" in persisted: + self.version = persisted["version"] + self._loaded_fields.add("version") + self._ledger_id = getattr(self, "id", self._ledger_id) + new_key = self._teaql_entity_key() + if old_key != new_key: + self._entity_root.rekey(old_key, new_key) + def rollback_entity(): + for field, value in rollback_payload.items(): + setattr(self, field, value) + self._ledger_id = rollback_ledger_id + self._action = rollback_action + self._loaded_fields = rollback_loaded_fields + if old_key != new_key: + self._entity_root.rekey(new_key, old_key) + context.after_graph_rollback(rollback_entity) + if action != "Delete": self._action = "Update" - elif action == "Update": - self.version = result.get("version", getattr(self, "version", None)) cascade_relations = [] - cascade_relations.append((self._order_line_list, "update_product")) + cascade_relations.append(("order_line_list", self._order_line_list, "update_product")) if action != "Delete": - for children, updater in cascade_relations: - for child in children: + for relation_name, children, updater in cascade_relations: + for index, child in enumerate(children): + child._teaql_attach_root(self._entity_root) getattr(child, updater)(self) child.audit_as(self._comment) - await child.save(context) - return result + try: + await child._teaql_save_within_graph(context) + except CheckException as error: + prefix = ObjectLocation().property(relation_name).index(index) + raise CheckException([ + CheckResult( + violation.rule_id, + violation.location.prefixed_by(prefix), + violation.input_value, + violation.system_value, + violation.message, + ) + for violation in error.violations + ]) from error + def commit_entity(): + self._entity_root.clear_entity(new_key) + if getattr(self, "version", None) is not None: + self._entity_root.set_original_version(new_key, int(self.version)) + context.after_graph_commit(commit_entity) + return self def update_id(self, value): self.id = value self._loaded_fields.add("id") + self._entity_root.set(self._teaql_entity_key(), "id", Value.from_any(value)) return self def update_name(self, value): self.name = value self._loaded_fields.add("name") + self._entity_root.set(self._teaql_entity_key(), "name", Value.from_any(value)) return self def update_sku(self, value): self.sku = value self._loaded_fields.add("sku") + self._entity_root.set(self._teaql_entity_key(), "sku", Value.from_any(value)) return self def update_image_url(self, value): self.imageUrl = value self._loaded_fields.add("imageUrl") + self._entity_root.set(self._teaql_entity_key(), "image_url", Value.from_any(value)) return self def update_create_time(self, value): self.createTime = value self._loaded_fields.add("createTime") + self._entity_root.set(self._teaql_entity_key(), "create_time", Value.from_any(value)) return self def update_update_time(self, value): self.updateTime = value self._loaded_fields.add("updateTime") + self._entity_root.set(self._teaql_entity_key(), "update_time", Value.from_any(value)) return self def update_version(self, value): self.version = value self._loaded_fields.add("version") + self._entity_root.set(self._teaql_entity_key(), "version", Value.from_any(value)) return self def update_commerce_platform(self, value): self.commercePlatform = getattr(value, "id", value) if value else None self._loaded_fields.add("commercePlatform") + self._entity_root.set(self._teaql_entity_key(), "commerce_platform", Value.from_any(self.commercePlatform)) return self def order_line_list(self) -> list: + self._loaded_fields.add("order_line_list") return self._order_line_list \ No newline at end of file diff --git a/examples/order-management/python-lib-core/pyproject.toml b/examples/order-management/python-lib-core/pyproject.toml index 60a5859..adfd0b1 100644 --- a/examples/order-management/python-lib-core/pyproject.toml +++ b/examples/order-management/python-lib-core/pyproject.toml @@ -2,14 +2,14 @@ name = "order-management-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] [tool.setuptools.packages.find] where = ["."] -include = ["models*", "requests*", "teaql*"] +include = ["models*", "requests*"] [build-system] requires = ["setuptools>=42"] diff --git a/examples/order-management/python-lib-core/requests/commerce_platform_request.py b/examples/order-management/python-lib-core/requests/commerce_platform_request.py index 24628e9..1379bb4 100644 --- a/examples/order-management/python-lib-core/requests/commerce_platform_request.py +++ b/examples/order-management/python-lib-core/requests/commerce_platform_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.commerce_platform import CommercePlatform +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class CommercePlatformRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("CommercePlatform") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableCommercePlatformRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,30 +64,290 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "name", "create_time", "update_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_name(self): + self.query.project("name") + return self + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_update_time(self): + self.query.project("update_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_name_containing(self, val: str): self.query.and_filter(contain("name", val)) return self + def with_name_not_containing(self, val: str): + self.query.and_filter(not_contain("name", val)) + return self + + def with_name_starting_with(self, val: str): + self.query.and_filter(begin_with("name", val)) + return self + + def with_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("name", val)) + return self + + def with_name_ending_with(self, val: str): + self.query.and_filter(end_with("name", val)) + return self + + def with_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("name", val)) + return self + + def with_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("name", val)) + return self + def with_name_is(self, val: str): self.query.and_filter(eq("name", val)) return self + def with_name_is_not(self, val): + self.query.and_filter(ne("name", val)) + return self + + def with_name_in(self, *vals): + self.query.and_filter(in_list("name", list(vals))) + return self + + def with_name_not_in(self, *vals): + self.query.and_filter(not_in_list("name", list(vals))) + return self + + def with_name_greater_than(self, val): + self.query.and_filter(gt("name", val)) + return self + + def with_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("name", val)) + return self + + def with_name_less_than(self, val): + self.query.and_filter(lt("name", val)) + return self + + def with_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("name", val)) + return self + + def with_name_between(self, lower, upper): + self.query.and_filter(between(column("name"), value(lower), value(upper))) + return self + + def with_name_is_known(self): + self.query.and_filter(is_not_null(column("name"))) + return self + + def with_name_is_unknown(self): + self.query.and_filter(is_null(column("name"))) + return self def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_update_time_is(self, val): self.query.and_filter(eq("update_time", val)) return self + def with_update_time_is_not(self, val): + self.query.and_filter(ne("update_time", val)) + return self + + def with_update_time_in(self, *vals): + self.query.and_filter(in_list("update_time", list(vals))) + return self + + def with_update_time_not_in(self, *vals): + self.query.and_filter(not_in_list("update_time", list(vals))) + return self + + def with_update_time_greater_than(self, val): + self.query.and_filter(gt("update_time", val)) + return self + + def with_update_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("update_time", val)) + return self + + def with_update_time_less_than(self, val): + self.query.and_filter(lt("update_time", val)) + return self + + def with_update_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("update_time", val)) + return self + + def with_update_time_between(self, lower, upper): + self.query.and_filter(between(column("update_time"), value(lower), value(upper))) + return self + + def with_update_time_is_known(self): + self.query.and_filter(is_not_null(column("update_time"))) + return self + + def with_update_time_is_unknown(self): + self.query.and_filter(is_null(column("update_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -179,37 +474,477 @@ def select_order_search_preset_list(self): def select_order_search_preset_list_with(self, child_request): self.query.relation_query("order_search_preset_list", child_request.query) return self + def have_customers(self): + from requests.customer_request import CustomerRequest + return self.with_customer_list_matching(CustomerRequest()) + + def have_no_customers(self): + from requests.customer_request import CustomerRequest + return self.without_customer_list_matching(CustomerRequest()) + + def with_customer_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "Customer", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_customer_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "Customer", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def have_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.with_order_status_list_matching(OrderStatusRequest()) + + def have_no_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.without_order_status_list_matching(OrderStatusRequest()) + + def with_order_status_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "OrderStatus", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_order_status_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "OrderStatus", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def have_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.with_customer_order_list_matching(CustomerOrderRequest()) + + def have_no_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.without_customer_order_list_matching(CustomerOrderRequest()) + + def with_customer_order_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_customer_order_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def have_products(self): + from requests.product_request import ProductRequest + return self.with_product_list_matching(ProductRequest()) + + def have_no_products(self): + from requests.product_request import ProductRequest + return self.without_product_list_matching(ProductRequest()) + + def with_product_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "Product", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_product_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "Product", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def have_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.with_order_line_list_matching(OrderLineRequest()) + + def have_no_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.without_order_line_list_matching(OrderLineRequest()) + + def with_order_line_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_order_line_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def have_order_search_presets(self): + from requests.order_search_preset_request import OrderSearchPresetRequest + return self.with_order_search_preset_list_matching(OrderSearchPresetRequest()) + + def have_no_order_search_presets(self): + from requests.order_search_preset_request import OrderSearchPresetRequest + return self.without_order_search_preset_list_matching(OrderSearchPresetRequest()) + + def with_order_search_preset_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "OrderSearchPreset", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + + def without_order_search_preset_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "OrderSearchPreset", child_request.query)) + child_request.query._projection = ["commerce_platform"] + return self + def count_customers(self): + return self.count_customers_as("count_customers") + + def count_customers_as(self, alias: str): + from requests.customer_request import CustomerRequest + return self.count_customers_with(alias, CustomerRequest()) + + def count_customers_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("customer_list", alias, child_request.query, True) + return self + + + def count_order_statuses(self): + return self.count_order_statuses_as("count_order_statuses") + + def count_order_statuses_as(self, alias: str): + from requests.order_status_request import OrderStatusRequest + return self.count_order_statuses_with(alias, OrderStatusRequest()) + + def count_order_statuses_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + + def min_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.min_display_order_of_order_statuses_as( + "min_display_order_of_order_statuses", OrderStatusRequest()) + + def min_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("min", "display_order", "min_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def max_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.max_display_order_of_order_statuses_as( + "max_display_order_of_order_statuses", OrderStatusRequest()) + + def max_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("max", "display_order", "max_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def sum_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.sum_display_order_of_order_statuses_as( + "sum_display_order_of_order_statuses", OrderStatusRequest()) + + def sum_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "display_order", "sum_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def avg_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.avg_display_order_of_order_statuses_as( + "avg_display_order_of_order_statuses", OrderStatusRequest()) + + def avg_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "display_order", "avg_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def standardDeviation_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.standardDeviation_display_order_of_order_statuses_as( + "standardDeviation_display_order_of_order_statuses", OrderStatusRequest()) + + def standardDeviation_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "display_order", "standardDeviation_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.squareRootOfPopulationStandardDeviation_display_order_of_order_statuses_as( + "squareRootOfPopulationStandardDeviation_display_order_of_order_statuses", OrderStatusRequest()) + + def squareRootOfPopulationStandardDeviation_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "display_order", "squareRootOfPopulationStandardDeviation_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def sampleVariance_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.sampleVariance_display_order_of_order_statuses_as( + "sampleVariance_display_order_of_order_statuses", OrderStatusRequest()) + + def sampleVariance_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "display_order", "sampleVariance_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def samplePopulationVariance_display_order_of_order_statuses(self): + from requests.order_status_request import OrderStatusRequest + return self.samplePopulationVariance_display_order_of_order_statuses_as( + "samplePopulationVariance_display_order_of_order_statuses", OrderStatusRequest()) + + def samplePopulationVariance_display_order_of_order_statuses_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "display_order", "samplePopulationVariance_display_order") + self.query.relation_aggregate("order_status_list", alias, child_request.query, True) + return self + def count_customer_orders(self): + return self.count_customer_orders_as("count_customer_orders") + + def count_customer_orders_as(self, alias: str): + from requests.customer_order_request import CustomerOrderRequest + return self.count_customer_orders_with(alias, CustomerOrderRequest()) + + def count_customer_orders_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + + def min_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.min_total_amount_of_customer_orders_as( + "min_total_amount_of_customer_orders", CustomerOrderRequest()) + + def min_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("min", "total_amount", "min_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def max_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.max_total_amount_of_customer_orders_as( + "max_total_amount_of_customer_orders", CustomerOrderRequest()) + + def max_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("max", "total_amount", "max_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sum_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sum_total_amount_of_customer_orders_as( + "sum_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sum_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "total_amount", "sum_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def avg_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.avg_total_amount_of_customer_orders_as( + "avg_total_amount_of_customer_orders", CustomerOrderRequest()) + + def avg_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "total_amount", "avg_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def standardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.standardDeviation_total_amount_of_customer_orders_as( + "standardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def standardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "total_amount", "standardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as( + "squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sampleVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sampleVariance_total_amount_of_customer_orders_as( + "sampleVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sampleVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "total_amount", "sampleVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def samplePopulationVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.samplePopulationVariance_total_amount_of_customer_orders_as( + "samplePopulationVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def samplePopulationVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def count_products(self): + return self.count_products_as("count_products") + + def count_products_as(self, alias: str): + from requests.product_request import ProductRequest + return self.count_products_with(alias, ProductRequest()) + + def count_products_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("product_list", alias, child_request.query, True) + return self + + + def count_order_lines(self): + return self.count_order_lines_as("count_order_lines") + + def count_order_lines_as(self, alias: str): + from requests.order_line_request import OrderLineRequest + return self.count_order_lines_with(alias, OrderLineRequest()) + + def count_order_lines_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + + def min_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.min_quantity_of_order_lines_as( + "min_quantity_of_order_lines", OrderLineRequest()) + + def min_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("min", "quantity", "min_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def max_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.max_quantity_of_order_lines_as( + "max_quantity_of_order_lines", OrderLineRequest()) + + def max_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("max", "quantity", "max_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sum_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sum_quantity_of_order_lines_as( + "sum_quantity_of_order_lines", OrderLineRequest()) + + def sum_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "quantity", "sum_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def avg_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.avg_quantity_of_order_lines_as( + "avg_quantity_of_order_lines", OrderLineRequest()) + + def avg_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "quantity", "avg_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def standardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.standardDeviation_quantity_of_order_lines_as( + "standardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def standardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "quantity", "standardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as( + "squareRootOfPopulationStandardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sampleVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sampleVariance_quantity_of_order_lines_as( + "sampleVariance_quantity_of_order_lines", OrderLineRequest()) + + def sampleVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "quantity", "sampleVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def samplePopulationVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.samplePopulationVariance_quantity_of_order_lines_as( + "samplePopulationVariance_quantity_of_order_lines", OrderLineRequest()) + + def samplePopulationVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def count_order_search_presets(self): + return self.count_order_search_presets_as("count_order_search_presets") + + def count_order_search_presets_as(self, alias: str): + from requests.order_search_preset_request import OrderSearchPresetRequest + return self.count_order_search_presets_with(alias, OrderSearchPresetRequest()) + + def count_order_search_presets_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("order_search_preset_list", alias, child_request.query, True) + return self + + class ExecutableCommercePlatformRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> CommercePlatform: - return CommercePlatform() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> CommercePlatform: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("CommercePlatform", CommercePlatform()) + if not isinstance(entity, CommercePlatform): + raise TypeError("entity initializer returned an incompatible CommercePlatform") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[CommercePlatform]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (CommercePlatform(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[CommercePlatform]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(CommercePlatform(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [CommercePlatform(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield CommercePlatform(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/customer_order_request.py b/examples/order-management/python-lib-core/requests/customer_order_request.py index d5f17c5..f7cc51c 100644 --- a/examples/order-management/python-lib-core/requests/customer_order_request.py +++ b/examples/order-management/python-lib-core/requests/customer_order_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.customer_order import CustomerOrder +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class CustomerOrderRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("CustomerOrder") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableCustomerOrderRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,29 +64,329 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "order_number", "order_date", "total_amount", "status", "customer", "commerce_platform", "create_time", "update_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_order_number(self): + self.query.project("order_number") + return self + + def select_order_date(self): + self.query.project("order_date") + return self + + def select_total_amount(self): + self.query.project("total_amount") + return self + + + + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_update_time(self): + self.query.project("update_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_status_with(self, child_request): + self.query.project("status") + self.query.relation_query("status", child_request.query) + return self + def select_customer_with(self, child_request): + self.query.project("customer") + self.query.relation_query("customer", child_request.query) + return self + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_status_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("status"), "OrderStatus", child_request.query)) + return self + + def without_status_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("status"), "OrderStatus", child_request.query)) + return self + + def have_status(self): + self.query.and_filter(is_not_null(column("status"))) + return self + + def have_no_status(self): + self.query.and_filter(is_null(column("status"))) + return self + def with_customer_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("customer"), "Customer", child_request.query)) + return self + + def without_customer_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("customer"), "Customer", child_request.query)) + return self + + def have_customer(self): + self.query.and_filter(is_not_null(column("customer"))) + return self + + def have_no_customer(self): + self.query.and_filter(is_null(column("customer"))) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_order_number_containing(self, val: str): self.query.and_filter(contain("order_number", val)) return self + def with_order_number_not_containing(self, val: str): + self.query.and_filter(not_contain("order_number", val)) + return self + + def with_order_number_starting_with(self, val: str): + self.query.and_filter(begin_with("order_number", val)) + return self + + def with_order_number_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("order_number", val)) + return self + + def with_order_number_ending_with(self, val: str): + self.query.and_filter(end_with("order_number", val)) + return self + + def with_order_number_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("order_number", val)) + return self + + def with_order_number_sounding_like(self, val: str): + self.query.and_filter(sound_like("order_number", val)) + return self + def with_order_number_is(self, val: str): self.query.and_filter(eq("order_number", val)) return self + def with_order_number_is_not(self, val): + self.query.and_filter(ne("order_number", val)) + return self + + def with_order_number_in(self, *vals): + self.query.and_filter(in_list("order_number", list(vals))) + return self + + def with_order_number_not_in(self, *vals): + self.query.and_filter(not_in_list("order_number", list(vals))) + return self + + def with_order_number_greater_than(self, val): + self.query.and_filter(gt("order_number", val)) + return self + + def with_order_number_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("order_number", val)) + return self + + def with_order_number_less_than(self, val): + self.query.and_filter(lt("order_number", val)) + return self + + def with_order_number_less_than_or_equal_to(self, val): + self.query.and_filter(lte("order_number", val)) + return self + + def with_order_number_between(self, lower, upper): + self.query.and_filter(between(column("order_number"), value(lower), value(upper))) + return self + + def with_order_number_is_known(self): + self.query.and_filter(is_not_null(column("order_number"))) + return self + + def with_order_number_is_unknown(self): + self.query.and_filter(is_null(column("order_number"))) + return self def with_order_date_is(self, val): self.query.and_filter(eq("order_date", val)) return self + def with_order_date_is_not(self, val): + self.query.and_filter(ne("order_date", val)) + return self + + def with_order_date_in(self, *vals): + self.query.and_filter(in_list("order_date", list(vals))) + return self + + def with_order_date_not_in(self, *vals): + self.query.and_filter(not_in_list("order_date", list(vals))) + return self + + def with_order_date_greater_than(self, val): + self.query.and_filter(gt("order_date", val)) + return self + + def with_order_date_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("order_date", val)) + return self + + def with_order_date_less_than(self, val): + self.query.and_filter(lt("order_date", val)) + return self + + def with_order_date_less_than_or_equal_to(self, val): + self.query.and_filter(lte("order_date", val)) + return self + + def with_order_date_between(self, lower, upper): + self.query.and_filter(between(column("order_date"), value(lower), value(upper))) + return self + + def with_order_date_is_known(self): + self.query.and_filter(is_not_null(column("order_date"))) + return self + + def with_order_date_is_unknown(self): + self.query.and_filter(is_null(column("order_date"))) + return self + def with_total_amount_is(self, val): self.query.and_filter(eq("total_amount", val)) return self + def with_total_amount_is_not(self, val): + self.query.and_filter(ne("total_amount", val)) + return self + + def with_total_amount_in(self, *vals): + self.query.and_filter(in_list("total_amount", list(vals))) + return self + + def with_total_amount_not_in(self, *vals): + self.query.and_filter(not_in_list("total_amount", list(vals))) + return self + + def with_total_amount_greater_than(self, val): + self.query.and_filter(gt("total_amount", val)) + return self + + def with_total_amount_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("total_amount", val)) + return self + + def with_total_amount_less_than(self, val): + self.query.and_filter(lt("total_amount", val)) + return self + + def with_total_amount_less_than_or_equal_to(self, val): + self.query.and_filter(lte("total_amount", val)) + return self + + def with_total_amount_between(self, lower, upper): + self.query.and_filter(between(column("total_amount"), value(lower), value(upper))) + return self + + def with_total_amount_is_known(self): + self.query.and_filter(is_not_null(column("total_amount"))) + return self + + def with_total_amount_is_unknown(self): + self.query.and_filter(is_null(column("total_amount"))) + return self + def filter_by_status(self, val): self.query.and_filter(eq("status", val)) return self + def with_status_is_pending(self): + self.query.and_filter(eq("status", 1001)) + return self + def with_status_is_confirmed(self): + self.query.and_filter(eq("status", 1002)) + return self def filter_by_customer(self, val): self.query.and_filter(eq("customer", val)) @@ -65,14 +400,134 @@ def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_update_time_is(self, val): self.query.and_filter(eq("update_time", val)) return self + def with_update_time_is_not(self, val): + self.query.and_filter(ne("update_time", val)) + return self + + def with_update_time_in(self, *vals): + self.query.and_filter(in_list("update_time", list(vals))) + return self + + def with_update_time_not_in(self, *vals): + self.query.and_filter(not_in_list("update_time", list(vals))) + return self + + def with_update_time_greater_than(self, val): + self.query.and_filter(gt("update_time", val)) + return self + + def with_update_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("update_time", val)) + return self + + def with_update_time_less_than(self, val): + self.query.and_filter(lt("update_time", val)) + return self + + def with_update_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("update_time", val)) + return self + + def with_update_time_between(self, lower, upper): + self.query.and_filter(between(column("update_time"), value(lower), value(upper))) + return self + + def with_update_time_is_known(self): + self.query.and_filter(is_not_null(column("update_time"))) + return self + + def with_update_time_is_unknown(self): + self.query.and_filter(is_null(column("update_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -266,37 +721,200 @@ def select_order_line_list(self): def select_order_line_list_with(self, child_request): self.query.relation_query("order_line_list", child_request.query) return self + def have_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.with_order_line_list_matching(OrderLineRequest()) + + def have_no_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.without_order_line_list_matching(OrderLineRequest()) + + def with_order_line_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["customer_order"] + return self + + def without_order_line_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["customer_order"] + return self + def count_order_lines(self): + return self.count_order_lines_as("count_order_lines") + + def count_order_lines_as(self, alias: str): + from requests.order_line_request import OrderLineRequest + return self.count_order_lines_with(alias, OrderLineRequest()) + + def count_order_lines_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + + def min_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.min_quantity_of_order_lines_as( + "min_quantity_of_order_lines", OrderLineRequest()) + + def min_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("min", "quantity", "min_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def max_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.max_quantity_of_order_lines_as( + "max_quantity_of_order_lines", OrderLineRequest()) + + def max_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("max", "quantity", "max_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sum_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sum_quantity_of_order_lines_as( + "sum_quantity_of_order_lines", OrderLineRequest()) + + def sum_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "quantity", "sum_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def avg_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.avg_quantity_of_order_lines_as( + "avg_quantity_of_order_lines", OrderLineRequest()) + + def avg_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "quantity", "avg_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def standardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.standardDeviation_quantity_of_order_lines_as( + "standardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def standardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "quantity", "standardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as( + "squareRootOfPopulationStandardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sampleVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sampleVariance_quantity_of_order_lines_as( + "sampleVariance_quantity_of_order_lines", OrderLineRequest()) + + def sampleVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "quantity", "sampleVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def samplePopulationVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.samplePopulationVariance_quantity_of_order_lines_as( + "samplePopulationVariance_quantity_of_order_lines", OrderLineRequest()) + + def samplePopulationVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def facet_by_status_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "status", request.query, include_all_facets) + return self + + def facet_by_customer_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "customer", request.query, include_all_facets) + return self + + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableCustomerOrderRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> CustomerOrder: - return CustomerOrder() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> CustomerOrder: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("CustomerOrder", CustomerOrder()) + if not isinstance(entity, CustomerOrder): + raise TypeError("entity initializer returned an incompatible CustomerOrder") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[CustomerOrder]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (CustomerOrder(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[CustomerOrder]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(CustomerOrder(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [CustomerOrder(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield CustomerOrder(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/customer_request.py b/examples/order-management/python-lib-core/requests/customer_request.py index d9c4b03..02c8e18 100644 --- a/examples/order-management/python-lib-core/requests/customer_request.py +++ b/examples/order-management/python-lib-core/requests/customer_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.customer import Customer +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class CustomerRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("Customer") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableCustomerRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,25 +64,254 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "name", "email", "commerce_platform", "create_time", "update_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_name(self): + self.query.project("name") + return self + + def select_email(self): + self.query.project("email") + return self + + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_update_time(self): + self.query.project("update_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_name_containing(self, val: str): self.query.and_filter(contain("name", val)) return self + def with_name_not_containing(self, val: str): + self.query.and_filter(not_contain("name", val)) + return self + + def with_name_starting_with(self, val: str): + self.query.and_filter(begin_with("name", val)) + return self + + def with_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("name", val)) + return self + + def with_name_ending_with(self, val: str): + self.query.and_filter(end_with("name", val)) + return self + + def with_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("name", val)) + return self + + def with_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("name", val)) + return self + def with_name_is(self, val: str): self.query.and_filter(eq("name", val)) return self + def with_name_is_not(self, val): + self.query.and_filter(ne("name", val)) + return self + + def with_name_in(self, *vals): + self.query.and_filter(in_list("name", list(vals))) + return self + + def with_name_not_in(self, *vals): + self.query.and_filter(not_in_list("name", list(vals))) + return self + + def with_name_greater_than(self, val): + self.query.and_filter(gt("name", val)) + return self + + def with_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("name", val)) + return self + + def with_name_less_than(self, val): + self.query.and_filter(lt("name", val)) + return self + + def with_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("name", val)) + return self + + def with_name_between(self, lower, upper): + self.query.and_filter(between(column("name"), value(lower), value(upper))) + return self + + def with_name_is_known(self): + self.query.and_filter(is_not_null(column("name"))) + return self + + def with_name_is_unknown(self): + self.query.and_filter(is_null(column("name"))) + return self def with_email_containing(self, val: str): self.query.and_filter(contain("email", val)) return self + def with_email_not_containing(self, val: str): + self.query.and_filter(not_contain("email", val)) + return self + + def with_email_starting_with(self, val: str): + self.query.and_filter(begin_with("email", val)) + return self + + def with_email_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("email", val)) + return self + + def with_email_ending_with(self, val: str): + self.query.and_filter(end_with("email", val)) + return self + + def with_email_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("email", val)) + return self + + def with_email_sounding_like(self, val: str): + self.query.and_filter(sound_like("email", val)) + return self + def with_email_is(self, val: str): self.query.and_filter(eq("email", val)) return self + def with_email_is_not(self, val): + self.query.and_filter(ne("email", val)) + return self + + def with_email_in(self, *vals): + self.query.and_filter(in_list("email", list(vals))) + return self + + def with_email_not_in(self, *vals): + self.query.and_filter(not_in_list("email", list(vals))) + return self + + def with_email_greater_than(self, val): + self.query.and_filter(gt("email", val)) + return self + + def with_email_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("email", val)) + return self + + def with_email_less_than(self, val): + self.query.and_filter(lt("email", val)) + return self + + def with_email_less_than_or_equal_to(self, val): + self.query.and_filter(lte("email", val)) + return self + + def with_email_between(self, lower, upper): + self.query.and_filter(between(column("email"), value(lower), value(upper))) + return self + + def with_email_is_known(self): + self.query.and_filter(is_not_null(column("email"))) + return self + + def with_email_is_unknown(self): + self.query.and_filter(is_null(column("email"))) + return self def filter_by_commerce_platform(self, val): self.query.and_filter(eq("commerce_platform", val)) @@ -57,14 +321,134 @@ def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_update_time_is(self, val): self.query.and_filter(eq("update_time", val)) return self + def with_update_time_is_not(self, val): + self.query.and_filter(ne("update_time", val)) + return self + + def with_update_time_in(self, *vals): + self.query.and_filter(in_list("update_time", list(vals))) + return self + + def with_update_time_not_in(self, *vals): + self.query.and_filter(not_in_list("update_time", list(vals))) + return self + + def with_update_time_greater_than(self, val): + self.query.and_filter(gt("update_time", val)) + return self + + def with_update_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("update_time", val)) + return self + + def with_update_time_less_than(self, val): + self.query.and_filter(lt("update_time", val)) + return self + + def with_update_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("update_time", val)) + return self + + def with_update_time_between(self, lower, upper): + self.query.and_filter(between(column("update_time"), value(lower), value(upper))) + return self + + def with_update_time_is_known(self): + self.query.and_filter(is_not_null(column("update_time"))) + return self + + def with_update_time_is_unknown(self): + self.query.and_filter(is_null(column("update_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -179,37 +563,190 @@ def select_customer_order_list(self): def select_customer_order_list_with(self, child_request): self.query.relation_query("customer_order_list", child_request.query) return self + def have_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.with_customer_order_list_matching(CustomerOrderRequest()) + + def have_no_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.without_customer_order_list_matching(CustomerOrderRequest()) + + def with_customer_order_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["customer"] + return self + + def without_customer_order_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["customer"] + return self + def count_customer_orders(self): + return self.count_customer_orders_as("count_customer_orders") + + def count_customer_orders_as(self, alias: str): + from requests.customer_order_request import CustomerOrderRequest + return self.count_customer_orders_with(alias, CustomerOrderRequest()) + + def count_customer_orders_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + + def min_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.min_total_amount_of_customer_orders_as( + "min_total_amount_of_customer_orders", CustomerOrderRequest()) + + def min_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("min", "total_amount", "min_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def max_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.max_total_amount_of_customer_orders_as( + "max_total_amount_of_customer_orders", CustomerOrderRequest()) + + def max_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("max", "total_amount", "max_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sum_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sum_total_amount_of_customer_orders_as( + "sum_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sum_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "total_amount", "sum_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def avg_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.avg_total_amount_of_customer_orders_as( + "avg_total_amount_of_customer_orders", CustomerOrderRequest()) + + def avg_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "total_amount", "avg_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def standardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.standardDeviation_total_amount_of_customer_orders_as( + "standardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def standardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "total_amount", "standardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as( + "squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sampleVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sampleVariance_total_amount_of_customer_orders_as( + "sampleVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sampleVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "total_amount", "sampleVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def samplePopulationVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.samplePopulationVariance_total_amount_of_customer_orders_as( + "samplePopulationVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def samplePopulationVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableCustomerRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> Customer: - return Customer() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> Customer: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("Customer", Customer()) + if not isinstance(entity, Customer): + raise TypeError("entity initializer returned an incompatible Customer") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[Customer]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (Customer(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[Customer]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(Customer(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [Customer(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield Customer(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/order_line_request.py b/examples/order-management/python-lib-core/requests/order_line_request.py index 4262223..a779f5a 100644 --- a/examples/order-management/python-lib-core/requests/order_line_request.py +++ b/examples/order-management/python-lib-core/requests/order_line_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.order_line import OrderLine +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class OrderLineRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("OrderLine") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableOrderLineRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,10 +64,157 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "customer_order", "product", "product_name", "sku", "quantity", "commerce_platform", "create_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + + + def select_product_name(self): + self.query.project("product_name") + return self + + def select_sku(self): + self.query.project("sku") + return self + + def select_quantity(self): + self.query.project("quantity") + return self + + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_customer_order_with(self, child_request): + self.query.project("customer_order") + self.query.relation_query("customer_order", child_request.query) + return self + def select_product_with(self, child_request): + self.query.project("product") + self.query.relation_query("product", child_request.query) + return self + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_customer_order_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("customer_order"), "CustomerOrder", child_request.query)) + return self + + def without_customer_order_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("customer_order"), "CustomerOrder", child_request.query)) + return self + + def have_customer_order(self): + self.query.and_filter(is_not_null(column("customer_order"))) + return self + + def have_no_customer_order(self): + self.query.and_filter(is_null(column("customer_order"))) + return self + def with_product_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("product"), "Product", child_request.query)) + return self + + def without_product_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("product"), "Product", child_request.query)) + return self + + def have_product(self): + self.query.and_filter(is_not_null(column("product"))) + return self + + def have_no_product(self): + self.query.and_filter(is_null(column("product"))) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def filter_by_customer_order(self, val): self.query.and_filter(eq("customer_order", val)) return self @@ -45,22 +227,188 @@ def with_product_name_containing(self, val: str): self.query.and_filter(contain("product_name", val)) return self + def with_product_name_not_containing(self, val: str): + self.query.and_filter(not_contain("product_name", val)) + return self + + def with_product_name_starting_with(self, val: str): + self.query.and_filter(begin_with("product_name", val)) + return self + + def with_product_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("product_name", val)) + return self + + def with_product_name_ending_with(self, val: str): + self.query.and_filter(end_with("product_name", val)) + return self + + def with_product_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("product_name", val)) + return self + + def with_product_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("product_name", val)) + return self + def with_product_name_is(self, val: str): self.query.and_filter(eq("product_name", val)) return self + def with_product_name_is_not(self, val): + self.query.and_filter(ne("product_name", val)) + return self + + def with_product_name_in(self, *vals): + self.query.and_filter(in_list("product_name", list(vals))) + return self + + def with_product_name_not_in(self, *vals): + self.query.and_filter(not_in_list("product_name", list(vals))) + return self + + def with_product_name_greater_than(self, val): + self.query.and_filter(gt("product_name", val)) + return self + + def with_product_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("product_name", val)) + return self + + def with_product_name_less_than(self, val): + self.query.and_filter(lt("product_name", val)) + return self + + def with_product_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("product_name", val)) + return self + + def with_product_name_between(self, lower, upper): + self.query.and_filter(between(column("product_name"), value(lower), value(upper))) + return self + + def with_product_name_is_known(self): + self.query.and_filter(is_not_null(column("product_name"))) + return self + + def with_product_name_is_unknown(self): + self.query.and_filter(is_null(column("product_name"))) + return self def with_sku_containing(self, val: str): self.query.and_filter(contain("sku", val)) return self + def with_sku_not_containing(self, val: str): + self.query.and_filter(not_contain("sku", val)) + return self + + def with_sku_starting_with(self, val: str): + self.query.and_filter(begin_with("sku", val)) + return self + + def with_sku_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("sku", val)) + return self + + def with_sku_ending_with(self, val: str): + self.query.and_filter(end_with("sku", val)) + return self + + def with_sku_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("sku", val)) + return self + + def with_sku_sounding_like(self, val: str): + self.query.and_filter(sound_like("sku", val)) + return self + def with_sku_is(self, val: str): self.query.and_filter(eq("sku", val)) return self + def with_sku_is_not(self, val): + self.query.and_filter(ne("sku", val)) + return self + + def with_sku_in(self, *vals): + self.query.and_filter(in_list("sku", list(vals))) + return self + + def with_sku_not_in(self, *vals): + self.query.and_filter(not_in_list("sku", list(vals))) + return self + + def with_sku_greater_than(self, val): + self.query.and_filter(gt("sku", val)) + return self + + def with_sku_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("sku", val)) + return self + + def with_sku_less_than(self, val): + self.query.and_filter(lt("sku", val)) + return self + + def with_sku_less_than_or_equal_to(self, val): + self.query.and_filter(lte("sku", val)) + return self + + def with_sku_between(self, lower, upper): + self.query.and_filter(between(column("sku"), value(lower), value(upper))) + return self + + def with_sku_is_known(self): + self.query.and_filter(is_not_null(column("sku"))) + return self + + def with_sku_is_unknown(self): + self.query.and_filter(is_null(column("sku"))) + return self def with_quantity_is(self, val): self.query.and_filter(eq("quantity", val)) return self + def with_quantity_is_not(self, val): + self.query.and_filter(ne("quantity", val)) + return self + + def with_quantity_in(self, *vals): + self.query.and_filter(in_list("quantity", list(vals))) + return self + + def with_quantity_not_in(self, *vals): + self.query.and_filter(not_in_list("quantity", list(vals))) + return self + + def with_quantity_greater_than(self, val): + self.query.and_filter(gt("quantity", val)) + return self + + def with_quantity_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("quantity", val)) + return self + + def with_quantity_less_than(self, val): + self.query.and_filter(lt("quantity", val)) + return self + + def with_quantity_less_than_or_equal_to(self, val): + self.query.and_filter(lte("quantity", val)) + return self + + def with_quantity_between(self, lower, upper): + self.query.and_filter(between(column("quantity"), value(lower), value(upper))) + return self + + def with_quantity_is_known(self): + self.query.and_filter(is_not_null(column("quantity"))) + return self + + def with_quantity_is_unknown(self): + self.query.and_filter(is_null(column("quantity"))) + return self + def filter_by_commerce_platform(self, val): self.query.and_filter(eq("commerce_platform", val)) return self @@ -69,10 +417,90 @@ def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -244,37 +672,99 @@ def group_by_version(self): def group_by_version_as(self, ret_name: str): self.query.group_by("version") return self + def facet_by_customer_order_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "customer_order", request.query, include_all_facets) + return self + + def facet_by_product_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "product", request.query, include_all_facets) + return self + + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableOrderLineRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> OrderLine: - return OrderLine() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> OrderLine: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("OrderLine", OrderLine()) + if not isinstance(entity, OrderLine): + raise TypeError("entity initializer returned an incompatible OrderLine") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[OrderLine]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (OrderLine(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[OrderLine]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(OrderLine(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [OrderLine(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield OrderLine(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/order_search_preset_request.py b/examples/order-management/python-lib-core/requests/order_search_preset_request.py index 5fd7fc8..04e0468 100644 --- a/examples/order-management/python-lib-core/requests/order_search_preset_request.py +++ b/examples/order-management/python-lib-core/requests/order_search_preset_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.order_search_preset import OrderSearchPreset +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class OrderSearchPresetRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("OrderSearchPreset") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableOrderSearchPresetRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,41 +64,404 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "name", "filter_json", "request_id", "owner_user_id", "commerce_platform", "create_time", "update_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_name(self): + self.query.project("name") + return self + + def select_filter_json(self): + self.query.project("filter_json") + return self + + def select_request_id(self): + self.query.project("request_id") + return self + + def select_owner_user_id(self): + self.query.project("owner_user_id") + return self + + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_update_time(self): + self.query.project("update_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_name_containing(self, val: str): self.query.and_filter(contain("name", val)) return self + def with_name_not_containing(self, val: str): + self.query.and_filter(not_contain("name", val)) + return self + + def with_name_starting_with(self, val: str): + self.query.and_filter(begin_with("name", val)) + return self + + def with_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("name", val)) + return self + + def with_name_ending_with(self, val: str): + self.query.and_filter(end_with("name", val)) + return self + + def with_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("name", val)) + return self + + def with_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("name", val)) + return self + def with_name_is(self, val: str): self.query.and_filter(eq("name", val)) return self + def with_name_is_not(self, val): + self.query.and_filter(ne("name", val)) + return self + + def with_name_in(self, *vals): + self.query.and_filter(in_list("name", list(vals))) + return self + + def with_name_not_in(self, *vals): + self.query.and_filter(not_in_list("name", list(vals))) + return self + + def with_name_greater_than(self, val): + self.query.and_filter(gt("name", val)) + return self + + def with_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("name", val)) + return self + + def with_name_less_than(self, val): + self.query.and_filter(lt("name", val)) + return self + + def with_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("name", val)) + return self + + def with_name_between(self, lower, upper): + self.query.and_filter(between(column("name"), value(lower), value(upper))) + return self + + def with_name_is_known(self): + self.query.and_filter(is_not_null(column("name"))) + return self + + def with_name_is_unknown(self): + self.query.and_filter(is_null(column("name"))) + return self def with_filter_json_containing(self, val: str): self.query.and_filter(contain("filter_json", val)) return self + def with_filter_json_not_containing(self, val: str): + self.query.and_filter(not_contain("filter_json", val)) + return self + + def with_filter_json_starting_with(self, val: str): + self.query.and_filter(begin_with("filter_json", val)) + return self + + def with_filter_json_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("filter_json", val)) + return self + + def with_filter_json_ending_with(self, val: str): + self.query.and_filter(end_with("filter_json", val)) + return self + + def with_filter_json_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("filter_json", val)) + return self + + def with_filter_json_sounding_like(self, val: str): + self.query.and_filter(sound_like("filter_json", val)) + return self + def with_filter_json_is(self, val: str): self.query.and_filter(eq("filter_json", val)) return self + def with_filter_json_is_not(self, val): + self.query.and_filter(ne("filter_json", val)) + return self + + def with_filter_json_in(self, *vals): + self.query.and_filter(in_list("filter_json", list(vals))) + return self + + def with_filter_json_not_in(self, *vals): + self.query.and_filter(not_in_list("filter_json", list(vals))) + return self + + def with_filter_json_greater_than(self, val): + self.query.and_filter(gt("filter_json", val)) + return self + + def with_filter_json_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("filter_json", val)) + return self + + def with_filter_json_less_than(self, val): + self.query.and_filter(lt("filter_json", val)) + return self + + def with_filter_json_less_than_or_equal_to(self, val): + self.query.and_filter(lte("filter_json", val)) + return self + + def with_filter_json_between(self, lower, upper): + self.query.and_filter(between(column("filter_json"), value(lower), value(upper))) + return self + + def with_filter_json_is_known(self): + self.query.and_filter(is_not_null(column("filter_json"))) + return self + + def with_filter_json_is_unknown(self): + self.query.and_filter(is_null(column("filter_json"))) + return self def with_request_id_containing(self, val: str): self.query.and_filter(contain("request_id", val)) return self + def with_request_id_not_containing(self, val: str): + self.query.and_filter(not_contain("request_id", val)) + return self + + def with_request_id_starting_with(self, val: str): + self.query.and_filter(begin_with("request_id", val)) + return self + + def with_request_id_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("request_id", val)) + return self + + def with_request_id_ending_with(self, val: str): + self.query.and_filter(end_with("request_id", val)) + return self + + def with_request_id_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("request_id", val)) + return self + + def with_request_id_sounding_like(self, val: str): + self.query.and_filter(sound_like("request_id", val)) + return self + def with_request_id_is(self, val: str): self.query.and_filter(eq("request_id", val)) return self + def with_request_id_is_not(self, val): + self.query.and_filter(ne("request_id", val)) + return self + + def with_request_id_in(self, *vals): + self.query.and_filter(in_list("request_id", list(vals))) + return self + + def with_request_id_not_in(self, *vals): + self.query.and_filter(not_in_list("request_id", list(vals))) + return self + + def with_request_id_greater_than(self, val): + self.query.and_filter(gt("request_id", val)) + return self + + def with_request_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("request_id", val)) + return self + + def with_request_id_less_than(self, val): + self.query.and_filter(lt("request_id", val)) + return self + + def with_request_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("request_id", val)) + return self + + def with_request_id_between(self, lower, upper): + self.query.and_filter(between(column("request_id"), value(lower), value(upper))) + return self + + def with_request_id_is_known(self): + self.query.and_filter(is_not_null(column("request_id"))) + return self + + def with_request_id_is_unknown(self): + self.query.and_filter(is_null(column("request_id"))) + return self def with_owner_user_id_containing(self, val: str): self.query.and_filter(contain("owner_user_id", val)) return self + def with_owner_user_id_not_containing(self, val: str): + self.query.and_filter(not_contain("owner_user_id", val)) + return self + + def with_owner_user_id_starting_with(self, val: str): + self.query.and_filter(begin_with("owner_user_id", val)) + return self + + def with_owner_user_id_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("owner_user_id", val)) + return self + + def with_owner_user_id_ending_with(self, val: str): + self.query.and_filter(end_with("owner_user_id", val)) + return self + + def with_owner_user_id_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("owner_user_id", val)) + return self + + def with_owner_user_id_sounding_like(self, val: str): + self.query.and_filter(sound_like("owner_user_id", val)) + return self + def with_owner_user_id_is(self, val: str): self.query.and_filter(eq("owner_user_id", val)) return self + def with_owner_user_id_is_not(self, val): + self.query.and_filter(ne("owner_user_id", val)) + return self + + def with_owner_user_id_in(self, *vals): + self.query.and_filter(in_list("owner_user_id", list(vals))) + return self + + def with_owner_user_id_not_in(self, *vals): + self.query.and_filter(not_in_list("owner_user_id", list(vals))) + return self + + def with_owner_user_id_greater_than(self, val): + self.query.and_filter(gt("owner_user_id", val)) + return self + + def with_owner_user_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("owner_user_id", val)) + return self + + def with_owner_user_id_less_than(self, val): + self.query.and_filter(lt("owner_user_id", val)) + return self + + def with_owner_user_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("owner_user_id", val)) + return self + + def with_owner_user_id_between(self, lower, upper): + self.query.and_filter(between(column("owner_user_id"), value(lower), value(upper))) + return self + + def with_owner_user_id_is_known(self): + self.query.and_filter(is_not_null(column("owner_user_id"))) + return self + + def with_owner_user_id_is_unknown(self): + self.query.and_filter(is_null(column("owner_user_id"))) + return self def filter_by_commerce_platform(self, val): self.query.and_filter(eq("commerce_platform", val)) @@ -73,14 +471,134 @@ def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_update_time_is(self, val): self.query.and_filter(eq("update_time", val)) return self + def with_update_time_is_not(self, val): + self.query.and_filter(ne("update_time", val)) + return self + + def with_update_time_in(self, *vals): + self.query.and_filter(in_list("update_time", list(vals))) + return self + + def with_update_time_not_in(self, *vals): + self.query.and_filter(not_in_list("update_time", list(vals))) + return self + + def with_update_time_greater_than(self, val): + self.query.and_filter(gt("update_time", val)) + return self + + def with_update_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("update_time", val)) + return self + + def with_update_time_less_than(self, val): + self.query.and_filter(lt("update_time", val)) + return self + + def with_update_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("update_time", val)) + return self + + def with_update_time_between(self, lower, upper): + self.query.and_filter(between(column("update_time"), value(lower), value(upper))) + return self + + def with_update_time_is_known(self): + self.query.and_filter(is_not_null(column("update_time"))) + return self + + def with_update_time_is_unknown(self): + self.query.and_filter(is_null(column("update_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -218,37 +736,89 @@ def group_by_version(self): def group_by_version_as(self, ret_name: str): self.query.group_by("version") return self + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableOrderSearchPresetRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> OrderSearchPreset: - return OrderSearchPreset() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> OrderSearchPreset: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("OrderSearchPreset", OrderSearchPreset()) + if not isinstance(entity, OrderSearchPreset): + raise TypeError("entity initializer returned an incompatible OrderSearchPreset") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[OrderSearchPreset]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (OrderSearchPreset(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[OrderSearchPreset]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(OrderSearchPreset(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [OrderSearchPreset(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield OrderSearchPreset(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/order_status_request.py b/examples/order-management/python-lib-core/requests/order_status_request.py index ec99aa3..4278758 100644 --- a/examples/order-management/python-lib-core/requests/order_status_request.py +++ b/examples/order-management/python-lib-core/requests/order_status_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.order_status import OrderStatus +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class OrderStatusRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("OrderStatus") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableOrderStatusRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,38 +64,370 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "name", "code", "color", "display_order", "commerce_platform", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_name(self): + self.query.project("name") + return self + + def select_code(self): + self.query.project("code") + return self + + def select_color(self): + self.query.project("color") + return self + + def select_display_order(self): + self.query.project("display_order") + return self + + + def select_version(self): + self.query.project("version") + return self + + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_name_containing(self, val: str): self.query.and_filter(contain("name", val)) return self + def with_name_not_containing(self, val: str): + self.query.and_filter(not_contain("name", val)) + return self + + def with_name_starting_with(self, val: str): + self.query.and_filter(begin_with("name", val)) + return self + + def with_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("name", val)) + return self + + def with_name_ending_with(self, val: str): + self.query.and_filter(end_with("name", val)) + return self + + def with_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("name", val)) + return self + + def with_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("name", val)) + return self + def with_name_is(self, val: str): self.query.and_filter(eq("name", val)) return self + def with_name_is_not(self, val): + self.query.and_filter(ne("name", val)) + return self + + def with_name_in(self, *vals): + self.query.and_filter(in_list("name", list(vals))) + return self + + def with_name_not_in(self, *vals): + self.query.and_filter(not_in_list("name", list(vals))) + return self + + def with_name_greater_than(self, val): + self.query.and_filter(gt("name", val)) + return self + + def with_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("name", val)) + return self + + def with_name_less_than(self, val): + self.query.and_filter(lt("name", val)) + return self + + def with_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("name", val)) + return self + + def with_name_between(self, lower, upper): + self.query.and_filter(between(column("name"), value(lower), value(upper))) + return self + + def with_name_is_known(self): + self.query.and_filter(is_not_null(column("name"))) + return self + + def with_name_is_unknown(self): + self.query.and_filter(is_null(column("name"))) + return self def with_code_containing(self, val: str): self.query.and_filter(contain("code", val)) return self + def with_code_not_containing(self, val: str): + self.query.and_filter(not_contain("code", val)) + return self + + def with_code_starting_with(self, val: str): + self.query.and_filter(begin_with("code", val)) + return self + + def with_code_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("code", val)) + return self + + def with_code_ending_with(self, val: str): + self.query.and_filter(end_with("code", val)) + return self + + def with_code_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("code", val)) + return self + + def with_code_sounding_like(self, val: str): + self.query.and_filter(sound_like("code", val)) + return self + def with_code_is(self, val: str): self.query.and_filter(eq("code", val)) return self + def with_code_is_not(self, val): + self.query.and_filter(ne("code", val)) + return self + + def with_code_in(self, *vals): + self.query.and_filter(in_list("code", list(vals))) + return self + + def with_code_not_in(self, *vals): + self.query.and_filter(not_in_list("code", list(vals))) + return self + + def with_code_greater_than(self, val): + self.query.and_filter(gt("code", val)) + return self + + def with_code_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("code", val)) + return self + + def with_code_less_than(self, val): + self.query.and_filter(lt("code", val)) + return self + + def with_code_less_than_or_equal_to(self, val): + self.query.and_filter(lte("code", val)) + return self + + def with_code_between(self, lower, upper): + self.query.and_filter(between(column("code"), value(lower), value(upper))) + return self + + def with_code_is_known(self): + self.query.and_filter(is_not_null(column("code"))) + return self + + def with_code_is_unknown(self): + self.query.and_filter(is_null(column("code"))) + return self def with_color_containing(self, val: str): self.query.and_filter(contain("color", val)) return self + def with_color_not_containing(self, val: str): + self.query.and_filter(not_contain("color", val)) + return self + + def with_color_starting_with(self, val: str): + self.query.and_filter(begin_with("color", val)) + return self + + def with_color_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("color", val)) + return self + + def with_color_ending_with(self, val: str): + self.query.and_filter(end_with("color", val)) + return self + + def with_color_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("color", val)) + return self + + def with_color_sounding_like(self, val: str): + self.query.and_filter(sound_like("color", val)) + return self + def with_color_is(self, val: str): self.query.and_filter(eq("color", val)) return self + def with_color_is_not(self, val): + self.query.and_filter(ne("color", val)) + return self + + def with_color_in(self, *vals): + self.query.and_filter(in_list("color", list(vals))) + return self + + def with_color_not_in(self, *vals): + self.query.and_filter(not_in_list("color", list(vals))) + return self + + def with_color_greater_than(self, val): + self.query.and_filter(gt("color", val)) + return self + + def with_color_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("color", val)) + return self + + def with_color_less_than(self, val): + self.query.and_filter(lt("color", val)) + return self + + def with_color_less_than_or_equal_to(self, val): + self.query.and_filter(lte("color", val)) + return self + + def with_color_between(self, lower, upper): + self.query.and_filter(between(column("color"), value(lower), value(upper))) + return self + + def with_color_is_known(self): + self.query.and_filter(is_not_null(column("color"))) + return self + + def with_color_is_unknown(self): + self.query.and_filter(is_null(column("color"))) + return self def with_display_order_is(self, val): self.query.and_filter(eq("display_order", val)) return self + def with_display_order_is_not(self, val): + self.query.and_filter(ne("display_order", val)) + return self + + def with_display_order_in(self, *vals): + self.query.and_filter(in_list("display_order", list(vals))) + return self + + def with_display_order_not_in(self, *vals): + self.query.and_filter(not_in_list("display_order", list(vals))) + return self + + def with_display_order_greater_than(self, val): + self.query.and_filter(gt("display_order", val)) + return self + + def with_display_order_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("display_order", val)) + return self + + def with_display_order_less_than(self, val): + self.query.and_filter(lt("display_order", val)) + return self + + def with_display_order_less_than_or_equal_to(self, val): + self.query.and_filter(lte("display_order", val)) + return self + + def with_display_order_between(self, lower, upper): + self.query.and_filter(between(column("display_order"), value(lower), value(upper))) + return self + + def with_display_order_is_known(self): + self.query.and_filter(is_not_null(column("display_order"))) + return self + + def with_display_order_is_unknown(self): + self.query.and_filter(is_null(column("display_order"))) + return self + def filter_by_commerce_platform(self, val): self.query.and_filter(eq("commerce_platform", val)) return self @@ -69,6 +436,46 @@ def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -231,37 +638,190 @@ def select_customer_order_list(self): def select_customer_order_list_with(self, child_request): self.query.relation_query("customer_order_list", child_request.query) return self + def have_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.with_customer_order_list_matching(CustomerOrderRequest()) + + def have_no_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.without_customer_order_list_matching(CustomerOrderRequest()) + + def with_customer_order_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["status"] + return self + + def without_customer_order_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "CustomerOrder", child_request.query)) + child_request.query._projection = ["status"] + return self + def count_customer_orders(self): + return self.count_customer_orders_as("count_customer_orders") + + def count_customer_orders_as(self, alias: str): + from requests.customer_order_request import CustomerOrderRequest + return self.count_customer_orders_with(alias, CustomerOrderRequest()) + + def count_customer_orders_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + + def min_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.min_total_amount_of_customer_orders_as( + "min_total_amount_of_customer_orders", CustomerOrderRequest()) + + def min_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("min", "total_amount", "min_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def max_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.max_total_amount_of_customer_orders_as( + "max_total_amount_of_customer_orders", CustomerOrderRequest()) + + def max_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("max", "total_amount", "max_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sum_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sum_total_amount_of_customer_orders_as( + "sum_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sum_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "total_amount", "sum_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def avg_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.avg_total_amount_of_customer_orders_as( + "avg_total_amount_of_customer_orders", CustomerOrderRequest()) + + def avg_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "total_amount", "avg_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def standardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.standardDeviation_total_amount_of_customer_orders_as( + "standardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def standardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "total_amount", "standardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as( + "squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders", CustomerOrderRequest()) + + def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def sampleVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.sampleVariance_total_amount_of_customer_orders_as( + "sampleVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def sampleVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "total_amount", "sampleVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def samplePopulationVariance_total_amount_of_customer_orders(self): + from requests.customer_order_request import CustomerOrderRequest + return self.samplePopulationVariance_total_amount_of_customer_orders_as( + "samplePopulationVariance_total_amount_of_customer_orders", CustomerOrderRequest()) + + def samplePopulationVariance_total_amount_of_customer_orders_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregate("customer_order_list", alias, child_request.query, True) + return self + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableOrderStatusRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> OrderStatus: - return OrderStatus() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> OrderStatus: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("OrderStatus", OrderStatus()) + if not isinstance(entity, OrderStatus): + raise TypeError("entity initializer returned an incompatible OrderStatus") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[OrderStatus]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (OrderStatus(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[OrderStatus]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(OrderStatus(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [OrderStatus(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield OrderStatus(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/requests/product_request.py b/examples/order-management/python-lib-core/requests/product_request.py index 1a3ad6c..9ee97e0 100644 --- a/examples/order-management/python-lib-core/requests/product_request.py +++ b/examples/order-management/python-lib-core/requests/product_request.py @@ -1,13 +1,30 @@ from teaql.core.query import SelectQuery +from teaql.core.list import SmartList, TeaQLPage +from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest -from teaql.core.expr import eq, contain +from teaql.core.expr import ( + begin_with, between, column, contain, end_with, eq, gt, gte, + in_list, in_subquery, is_not_null, is_null, lt, lte, ne, not_begin_with, + not_contain, not_end_with, not_in_list, not_in_subquery, value, + sound_like, +) from models.product import Product +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class ProductRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("Product") self._purpose = None self._comment = None + self.query.and_filter(gte("version", 1)) + if minimal: + self.select_id() + self.select_version() + else: + self.select_self_fields() def comment(self, c: str): self.query.comment(c) @@ -15,12 +32,30 @@ def comment(self, c: str): return self def purpose(self, p: str): - if not self._comment or not self._comment.strip(): - raise ValueError("purpose() requires a non-empty comment() set earlier on the request") self.query.purpose(p) self._purpose = p return ExecutableProductRequest(self) + def optimize_for_continuous_page_fetch(self): + self.query.optimize_for_continuous_page_fetch() + return self + + def optimize_for_continuous_page_fetch_with(self, namespace: str, ttl_seconds: int): + self.query.optimize_for_continuous_page_fetch_with(namespace, ttl_seconds) + return self + + def optimize_pagination_with_id_set(self): + self.query.optimize_pagination_with_id_set() + return self + + def optimize_pagination_with_id_set_config(self, namespace: str, ttl_seconds: int, max_ids: int): + self.query.optimize_pagination_with_id_set_config(namespace, ttl_seconds, max_ids) + return self + + def top_n_probe_parent_threshold(self, threshold: int): + self.query.top_n_probe_parent_threshold(threshold) + return self + def limit(self, n: int): self.query.limit(n) return self @@ -29,33 +64,329 @@ def offset(self, n: int): self.query.offset(n) return self + def with_deleted_rows(self): + self.query._filters = [ + expression for expression in self.query._filters + if expression.get("field") != "version" + ] + return self + + def deleted_rows_only(self): + self.with_deleted_rows() + self.query.and_filter(lte("version", -1)) + return self + + def select_self_fields(self): + self.query.project("id", "name", "sku", "image_url", "commerce_platform", "create_time", "update_time", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + def select_name(self): + self.query.project("name") + return self + + def select_sku(self): + self.query.project("sku") + return self + + def select_image_url(self): + self.query.project("image_url") + return self + + + def select_create_time(self): + self.query.project("create_time") + return self + + def select_update_time(self): + self.query.project("update_time") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_commerce_platform_with(self, child_request): + self.query.project("commerce_platform") + self.query.relation_query("commerce_platform", child_request.query) + return self + def with_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def without_commerce_platform_matching(self, child_request): + child_request.query._projection = ["id"] + self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) + return self + + def have_commerce_platform(self): + self.query.and_filter(is_not_null(column("commerce_platform"))) + return self + + def have_no_commerce_platform(self): + self.query.and_filter(is_null(column("commerce_platform"))) + return self + def with_id_is(self, val): self.query.and_filter(eq("id", val)) return self + def with_id_is_not(self, val): + self.query.and_filter(ne("id", val)) + return self + + def with_id_in(self, *vals): + self.query.and_filter(in_list("id", list(vals))) + return self + + def with_id_not_in(self, *vals): + self.query.and_filter(not_in_list("id", list(vals))) + return self + + def with_id_greater_than(self, val): + self.query.and_filter(gt("id", val)) + return self + + def with_id_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("id", val)) + return self + + def with_id_less_than(self, val): + self.query.and_filter(lt("id", val)) + return self + + def with_id_less_than_or_equal_to(self, val): + self.query.and_filter(lte("id", val)) + return self + + def with_id_between(self, lower, upper): + self.query.and_filter(between(column("id"), value(lower), value(upper))) + return self + + def with_id_is_known(self): + self.query.and_filter(is_not_null(column("id"))) + return self + + def with_id_is_unknown(self): + self.query.and_filter(is_null(column("id"))) + return self + def with_name_containing(self, val: str): self.query.and_filter(contain("name", val)) return self + def with_name_not_containing(self, val: str): + self.query.and_filter(not_contain("name", val)) + return self + + def with_name_starting_with(self, val: str): + self.query.and_filter(begin_with("name", val)) + return self + + def with_name_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("name", val)) + return self + + def with_name_ending_with(self, val: str): + self.query.and_filter(end_with("name", val)) + return self + + def with_name_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("name", val)) + return self + + def with_name_sounding_like(self, val: str): + self.query.and_filter(sound_like("name", val)) + return self + def with_name_is(self, val: str): self.query.and_filter(eq("name", val)) return self + def with_name_is_not(self, val): + self.query.and_filter(ne("name", val)) + return self + + def with_name_in(self, *vals): + self.query.and_filter(in_list("name", list(vals))) + return self + + def with_name_not_in(self, *vals): + self.query.and_filter(not_in_list("name", list(vals))) + return self + + def with_name_greater_than(self, val): + self.query.and_filter(gt("name", val)) + return self + + def with_name_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("name", val)) + return self + + def with_name_less_than(self, val): + self.query.and_filter(lt("name", val)) + return self + + def with_name_less_than_or_equal_to(self, val): + self.query.and_filter(lte("name", val)) + return self + + def with_name_between(self, lower, upper): + self.query.and_filter(between(column("name"), value(lower), value(upper))) + return self + + def with_name_is_known(self): + self.query.and_filter(is_not_null(column("name"))) + return self + + def with_name_is_unknown(self): + self.query.and_filter(is_null(column("name"))) + return self def with_sku_containing(self, val: str): self.query.and_filter(contain("sku", val)) return self + def with_sku_not_containing(self, val: str): + self.query.and_filter(not_contain("sku", val)) + return self + + def with_sku_starting_with(self, val: str): + self.query.and_filter(begin_with("sku", val)) + return self + + def with_sku_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("sku", val)) + return self + + def with_sku_ending_with(self, val: str): + self.query.and_filter(end_with("sku", val)) + return self + + def with_sku_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("sku", val)) + return self + + def with_sku_sounding_like(self, val: str): + self.query.and_filter(sound_like("sku", val)) + return self + def with_sku_is(self, val: str): self.query.and_filter(eq("sku", val)) return self + def with_sku_is_not(self, val): + self.query.and_filter(ne("sku", val)) + return self + + def with_sku_in(self, *vals): + self.query.and_filter(in_list("sku", list(vals))) + return self + + def with_sku_not_in(self, *vals): + self.query.and_filter(not_in_list("sku", list(vals))) + return self + + def with_sku_greater_than(self, val): + self.query.and_filter(gt("sku", val)) + return self + + def with_sku_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("sku", val)) + return self + + def with_sku_less_than(self, val): + self.query.and_filter(lt("sku", val)) + return self + + def with_sku_less_than_or_equal_to(self, val): + self.query.and_filter(lte("sku", val)) + return self + + def with_sku_between(self, lower, upper): + self.query.and_filter(between(column("sku"), value(lower), value(upper))) + return self + + def with_sku_is_known(self): + self.query.and_filter(is_not_null(column("sku"))) + return self + + def with_sku_is_unknown(self): + self.query.and_filter(is_null(column("sku"))) + return self def with_image_url_containing(self, val: str): self.query.and_filter(contain("image_url", val)) return self + def with_image_url_not_containing(self, val: str): + self.query.and_filter(not_contain("image_url", val)) + return self + + def with_image_url_starting_with(self, val: str): + self.query.and_filter(begin_with("image_url", val)) + return self + + def with_image_url_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("image_url", val)) + return self + + def with_image_url_ending_with(self, val: str): + self.query.and_filter(end_with("image_url", val)) + return self + + def with_image_url_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("image_url", val)) + return self + + def with_image_url_sounding_like(self, val: str): + self.query.and_filter(sound_like("image_url", val)) + return self + def with_image_url_is(self, val: str): self.query.and_filter(eq("image_url", val)) return self + def with_image_url_is_not(self, val): + self.query.and_filter(ne("image_url", val)) + return self + + def with_image_url_in(self, *vals): + self.query.and_filter(in_list("image_url", list(vals))) + return self + + def with_image_url_not_in(self, *vals): + self.query.and_filter(not_in_list("image_url", list(vals))) + return self + + def with_image_url_greater_than(self, val): + self.query.and_filter(gt("image_url", val)) + return self + + def with_image_url_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("image_url", val)) + return self + + def with_image_url_less_than(self, val): + self.query.and_filter(lt("image_url", val)) + return self + + def with_image_url_less_than_or_equal_to(self, val): + self.query.and_filter(lte("image_url", val)) + return self + + def with_image_url_between(self, lower, upper): + self.query.and_filter(between(column("image_url"), value(lower), value(upper))) + return self + + def with_image_url_is_known(self): + self.query.and_filter(is_not_null(column("image_url"))) + return self + + def with_image_url_is_unknown(self): + self.query.and_filter(is_null(column("image_url"))) + return self def filter_by_commerce_platform(self, val): self.query.and_filter(eq("commerce_platform", val)) @@ -65,14 +396,134 @@ def with_create_time_is(self, val): self.query.and_filter(eq("create_time", val)) return self + def with_create_time_is_not(self, val): + self.query.and_filter(ne("create_time", val)) + return self + + def with_create_time_in(self, *vals): + self.query.and_filter(in_list("create_time", list(vals))) + return self + + def with_create_time_not_in(self, *vals): + self.query.and_filter(not_in_list("create_time", list(vals))) + return self + + def with_create_time_greater_than(self, val): + self.query.and_filter(gt("create_time", val)) + return self + + def with_create_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("create_time", val)) + return self + + def with_create_time_less_than(self, val): + self.query.and_filter(lt("create_time", val)) + return self + + def with_create_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("create_time", val)) + return self + + def with_create_time_between(self, lower, upper): + self.query.and_filter(between(column("create_time"), value(lower), value(upper))) + return self + + def with_create_time_is_known(self): + self.query.and_filter(is_not_null(column("create_time"))) + return self + + def with_create_time_is_unknown(self): + self.query.and_filter(is_null(column("create_time"))) + return self + def with_update_time_is(self, val): self.query.and_filter(eq("update_time", val)) return self + def with_update_time_is_not(self, val): + self.query.and_filter(ne("update_time", val)) + return self + + def with_update_time_in(self, *vals): + self.query.and_filter(in_list("update_time", list(vals))) + return self + + def with_update_time_not_in(self, *vals): + self.query.and_filter(not_in_list("update_time", list(vals))) + return self + + def with_update_time_greater_than(self, val): + self.query.and_filter(gt("update_time", val)) + return self + + def with_update_time_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("update_time", val)) + return self + + def with_update_time_less_than(self, val): + self.query.and_filter(lt("update_time", val)) + return self + + def with_update_time_less_than_or_equal_to(self, val): + self.query.and_filter(lte("update_time", val)) + return self + + def with_update_time_between(self, lower, upper): + self.query.and_filter(between(column("update_time"), value(lower), value(upper))) + return self + + def with_update_time_is_known(self): + self.query.and_filter(is_not_null(column("update_time"))) + return self + + def with_update_time_is_unknown(self): + self.query.and_filter(is_null(column("update_time"))) + return self + def with_version_is(self, val): self.query.and_filter(eq("version", val)) return self + def with_version_is_not(self, val): + self.query.and_filter(ne("version", val)) + return self + + def with_version_in(self, *vals): + self.query.and_filter(in_list("version", list(vals))) + return self + + def with_version_not_in(self, *vals): + self.query.and_filter(not_in_list("version", list(vals))) + return self + + def with_version_greater_than(self, val): + self.query.and_filter(gt("version", val)) + return self + + def with_version_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("version", val)) + return self + + def with_version_less_than(self, val): + self.query.and_filter(lt("version", val)) + return self + + def with_version_less_than_or_equal_to(self, val): + self.query.and_filter(lte("version", val)) + return self + + def with_version_between(self, lower, upper): + self.query.and_filter(between(column("version"), value(lower), value(upper))) + return self + + def with_version_is_known(self): + self.query.and_filter(is_not_null(column("version"))) + return self + + def with_version_is_unknown(self): + self.query.and_filter(is_null(column("version"))) + return self + def order_by_id_ascending(self): self.query.order_by("id", "asc") return self @@ -202,37 +653,190 @@ def select_order_line_list(self): def select_order_line_list_with(self, child_request): self.query.relation_query("order_line_list", child_request.query) return self + def have_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.with_order_line_list_matching(OrderLineRequest()) + + def have_no_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.without_order_line_list_matching(OrderLineRequest()) + + def with_order_line_list_matching(self, child_request): + self.query.and_filter(in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["product"] + return self + + def without_order_line_list_matching(self, child_request): + self.query.and_filter(not_in_subquery(column("id"), "OrderLine", child_request.query)) + child_request.query._projection = ["product"] + return self + def count_order_lines(self): + return self.count_order_lines_as("count_order_lines") + + def count_order_lines_as(self, alias: str): + from requests.order_line_request import OrderLineRequest + return self.count_order_lines_with(alias, OrderLineRequest()) + + def count_order_lines_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + + def min_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.min_quantity_of_order_lines_as( + "min_quantity_of_order_lines", OrderLineRequest()) + + def min_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("min", "quantity", "min_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def max_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.max_quantity_of_order_lines_as( + "max_quantity_of_order_lines", OrderLineRequest()) + + def max_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("max", "quantity", "max_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sum_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sum_quantity_of_order_lines_as( + "sum_quantity_of_order_lines", OrderLineRequest()) + + def sum_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("sum", "quantity", "sum_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def avg_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.avg_quantity_of_order_lines_as( + "avg_quantity_of_order_lines", OrderLineRequest()) + + def avg_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("avg", "quantity", "avg_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def standardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.standardDeviation_quantity_of_order_lines_as( + "standardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def standardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev", "quantity", "standardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as( + "squareRootOfPopulationStandardDeviation_quantity_of_order_lines", OrderLineRequest()) + + def squareRootOfPopulationStandardDeviation_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("stddev_pop", "quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def sampleVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.sampleVariance_quantity_of_order_lines_as( + "sampleVariance_quantity_of_order_lines", OrderLineRequest()) + + def sampleVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_samp", "quantity", "sampleVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def samplePopulationVariance_quantity_of_order_lines(self): + from requests.order_line_request import OrderLineRequest + return self.samplePopulationVariance_quantity_of_order_lines_as( + "samplePopulationVariance_quantity_of_order_lines", OrderLineRequest()) + + def samplePopulationVariance_quantity_of_order_lines_as(self, alias: str, child_request): + child_request.query.aggregate("var_pop", "quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregate("order_line_list", alias, child_request.query, True) + return self + def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "commerce_platform", request.query, include_all_facets) + return self + class ExecutableProductRequest: def __init__(self, request): self._request = request - def new_entity(self, context) -> Product: - return Product() + def comment(self, c: str): + self._request.comment(c) + return self - async def execute_for_list(self, context): + def new_entity(self, context) -> Product: + request = self._request + if not request._comment or not request._comment.strip() or not request._purpose or not request._purpose.strip(): + raise ValueError("Security audit failure: non-empty comment() and purpose() are required before new_entity()") + entity = context.initialize_entity("Product", Product()) + if not isinstance(entity, Product): + raise TypeError("entity initializer returned an incompatible Product") + return entity + + async def execute_for_result(self, context): self = self._request - if not self._purpose or not self._comment: - raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_list()") + if not self._purpose or not self._purpose.strip() or not self._comment or not self._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_rows()") service = context.require_resource("dataService") - req = QueryRequest(self.query) - res = await service.query(context, req) - - result = {"data": res.rows} - return result + req = QueryRequest(context.prepare_query(self.query)) + return await service.query(context, req) + + async def execute_for_rows(self, context): + return (await self.execute_for_result(context)).rows + + async def execute_for_list(self, context) -> SmartList[Product]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (Product(_entity_root=query_root, **row) for row in result.rows), + facets=result.facets) + + async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[Product]: + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise ValueError("Security audit failure: comment() and purpose() must be called before execute_for_page()") + request.query.offset(offset).limit(limit) + authorized = context.prepare_query(request.query) + service = context.require_resource("dataService") + alias = "__teaql_total" + if authorized.id_set_pagination is not None: + row_result = await service.query(context, QueryRequest(authorized)) + retained_count, accuracy = context.id_set_count() + if accuracy == "EXACT": + total_count = retained_count + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + else: + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias))) + if not count_result.rows or not isinstance(count_result.rows[0].get(alias), (int, float)): + raise RuntimeError("dataService did not return an exact page count") + total_count = int(count_result.rows[0][alias]) + row_result = await service.query(context, QueryRequest(authorized)) + query_root = EntityRoot() + data = SmartList(Product(_entity_root=query_root, **row) for row in row_result.rows) + return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) async def execute_for_one(self, context): self._request.limit(1) - res = await self.execute_for_list(context) - if res["data"]: - return res["data"][0] - return None - - async def execute_entities_for_list(self, context): - res = await self.execute_for_list(context) - return [Product(**row) for row in res["data"]] - - async def execute_entity_for_one(self, context): - self._request.limit(1) - entities = await self.execute_entities_for_list(context) - return entities[0] if entities else None \ No newline at end of file + entities = await self.execute_for_list(context) + return entities[0] if entities else None + + async def execute_for_stream(self, context, chunk_size: int = 1000): + """Yield entity chunks lazily from the provider cursor.""" + request = self._request + if not request._purpose or not request._purpose.strip() or not request._comment or not request._comment.strip(): + raise Exception("Security audit failure: comment() and purpose() must be called before execute_for_stream()") + service = context.require_resource("dataService") + if not hasattr(service, "query_stream"): + raise RuntimeError("dataService does not implement query_stream") + query_root = EntityRoot() + async for chunk in service.query_stream(context, QueryRequest(request.query), chunk_size): + for row in chunk.rows: + yield Product(_entity_root=query_root, **row) diff --git a/examples/order-management/python-lib-core/runtime_module.py b/examples/order-management/python-lib-core/runtime_module.py new file mode 100644 index 0000000..4979871 --- /dev/null +++ b/examples/order-management/python-lib-core/runtime_module.py @@ -0,0 +1,460 @@ +import asyncio +from datetime import datetime, timezone +from teaql.runtime import CheckResult, ContextEntityRef, JsonFieldNamingProfile, ObjectLocation, RuntimeModule, create_wire_entity_metadata +from teaql.core.meta import EntityDescriptor, PropertyDescriptor, RelationDescriptor +from teaql.core.value import DataType +from Q import Q +from teaql.core.value import Value +try: + from teaql.core.graph import GraphNode +except ImportError: + class GraphNode: + def __init__(self, entity): + self.entity, self.fields = entity, {} + def set(self, field, value): + self.fields[field] = value + return self +from models.commerce_platform import CommercePlatform +from models.customer import Customer +from models.order_status import OrderStatus +from models.customer_order import CustomerOrder +from models.product import Product +from models.order_line import OrderLine +from models.order_search_preset import OrderSearchPreset + +def _teaql_is_null(value): + return value.is_null() if hasattr(value, "is_null") else value is None + +def _teaql_raw(value): + return value.val if hasattr(value, "val") else value + +def _teaql_entity_id(value): + value = _teaql_raw(value) + if hasattr(value, "id"): + return value.id + if isinstance(value, dict): + return value.get("id") + return value + +class _CommercePlatformChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("CommercePlatform", "create_time", "clock", "graphClock") + + if operation == "insert" and ("update_time" not in record or _teaql_is_null(record["update_time"])): + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("CommercePlatform", "update_time", "clock", "graphClock") + if operation == "insert" or operation == "update": + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("CommercePlatform", "update_time", "clock", "graphClock") + + + if (operation == "insert" and "name" not in record) or ("name" in record and _teaql_is_null(record["name"])): + results.append(CheckResult("required", ObjectLocation().property("name"))) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 100)) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + if (operation == "insert" and "update_time" not in record) or ("update_time" in record and _teaql_is_null(record["update_time"])): + results.append(CheckResult("required", ObjectLocation().property("update_time"))) + + + +class _CustomerChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("Customer", "create_time", "clock", "graphClock") + + if operation == "insert" and ("update_time" not in record or _teaql_is_null(record["update_time"])): + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("Customer", "update_time", "clock", "graphClock") + if operation == "insert" or operation == "update": + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("Customer", "update_time", "clock", "graphClock") + + + if (operation == "insert" and "name" not in record) or ("name" in record and _teaql_is_null(record["name"])): + results.append(CheckResult("required", ObjectLocation().property("name"))) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 100)) + + if (operation == "insert" and "email" not in record) or ("email" in record and _teaql_is_null(record["email"])): + results.append(CheckResult("required", ObjectLocation().property("email"))) + if "email" in record and _teaql_raw(record["email"]) is not None and len(_teaql_raw(record["email"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("email"), _teaql_raw(record["email"]), 100)) + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + if (operation == "insert" and "update_time" not in record) or ("update_time" in record and _teaql_is_null(record["update_time"])): + results.append(CheckResult("required", ObjectLocation().property("update_time"))) + + + +class _OrderStatusChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if (operation == "insert" and "name" not in record) or ("name" in record and _teaql_is_null(record["name"])): + results.append(CheckResult("required", ObjectLocation().property("name"))) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 100)) + + if (operation == "insert" and "code" not in record) or ("code" in record and _teaql_is_null(record["code"])): + results.append(CheckResult("required", ObjectLocation().property("code"))) + if "code" in record and _teaql_raw(record["code"]) is not None and len(_teaql_raw(record["code"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("code"), _teaql_raw(record["code"]), 100)) + + if "color" in record and _teaql_raw(record["color"]) is not None and len(_teaql_raw(record["color"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("color"), _teaql_raw(record["color"]), 100)) + + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + + +class _CustomerOrderChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("CustomerOrder", "create_time", "clock", "graphClock") + + if operation == "insert" and ("update_time" not in record or _teaql_is_null(record["update_time"])): + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("CustomerOrder", "update_time", "clock", "graphClock") + if operation == "insert" or operation == "update": + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("CustomerOrder", "update_time", "clock", "graphClock") + + + if (operation == "insert" and "order_number" not in record) or ("order_number" in record and _teaql_is_null(record["order_number"])): + results.append(CheckResult("required", ObjectLocation().property("order_number"))) + if "order_number" in record and _teaql_raw(record["order_number"]) is not None and len(_teaql_raw(record["order_number"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("order_number"), _teaql_raw(record["order_number"]), 100)) + + if (operation == "insert" and "order_date" not in record) or ("order_date" in record and _teaql_is_null(record["order_date"])): + results.append(CheckResult("required", ObjectLocation().property("order_date"))) + + if (operation == "insert" and "total_amount" not in record) or ("total_amount" in record and _teaql_is_null(record["total_amount"])): + results.append(CheckResult("required", ObjectLocation().property("total_amount"))) + + if (operation == "insert" and "status" not in record) or ("status" in record and _teaql_is_null(record["status"])): + results.append(CheckResult("required", ObjectLocation().property("status"))) + + if (operation == "insert" and "customer" not in record) or ("customer" in record and _teaql_is_null(record["customer"])): + results.append(CheckResult("required", ObjectLocation().property("customer"))) + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + if (operation == "insert" and "update_time" not in record) or ("update_time" in record and _teaql_is_null(record["update_time"])): + results.append(CheckResult("required", ObjectLocation().property("update_time"))) + + + +class _ProductChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("Product", "create_time", "clock", "graphClock") + + if operation == "insert" and ("update_time" not in record or _teaql_is_null(record["update_time"])): + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("Product", "update_time", "clock", "graphClock") + if operation == "insert" or operation == "update": + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("Product", "update_time", "clock", "graphClock") + + + if (operation == "insert" and "name" not in record) or ("name" in record and _teaql_is_null(record["name"])): + results.append(CheckResult("required", ObjectLocation().property("name"))) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 100)) + + if (operation == "insert" and "sku" not in record) or ("sku" in record and _teaql_is_null(record["sku"])): + results.append(CheckResult("required", ObjectLocation().property("sku"))) + if "sku" in record and _teaql_raw(record["sku"]) is not None and len(_teaql_raw(record["sku"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("sku"), _teaql_raw(record["sku"]), 100)) + + if "image_url" in record and _teaql_raw(record["image_url"]) is not None and len(_teaql_raw(record["image_url"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("image_url"), _teaql_raw(record["image_url"]), 100)) + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + if (operation == "insert" and "update_time" not in record) or ("update_time" in record and _teaql_is_null(record["update_time"])): + results.append(CheckResult("required", ObjectLocation().property("update_time"))) + + + +class _OrderLineChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("OrderLine", "create_time", "clock", "graphClock") + + + if (operation == "insert" and "customer_order" not in record) or ("customer_order" in record and _teaql_is_null(record["customer_order"])): + results.append(CheckResult("required", ObjectLocation().property("customer_order"))) + + if (operation == "insert" and "product" not in record) or ("product" in record and _teaql_is_null(record["product"])): + results.append(CheckResult("required", ObjectLocation().property("product"))) + + if (operation == "insert" and "product_name" not in record) or ("product_name" in record and _teaql_is_null(record["product_name"])): + results.append(CheckResult("required", ObjectLocation().property("product_name"))) + if "product_name" in record and _teaql_raw(record["product_name"]) is not None and len(_teaql_raw(record["product_name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("product_name"), _teaql_raw(record["product_name"]), 100)) + + if (operation == "insert" and "sku" not in record) or ("sku" in record and _teaql_is_null(record["sku"])): + results.append(CheckResult("required", ObjectLocation().property("sku"))) + if "sku" in record and _teaql_raw(record["sku"]) is not None and len(_teaql_raw(record["sku"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("sku"), _teaql_raw(record["sku"]), 100)) + + if (operation == "insert" and "quantity" not in record) or ("quantity" in record and _teaql_is_null(record["quantity"])): + results.append(CheckResult("required", ObjectLocation().property("quantity"))) + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + + +class _OrderSearchPresetChecker: + def check_and_fix(self, context, record, location, results): + operation = context.get_resource("fix_operation") + now = context.get_resource("fix_time") + if operation == "insert" and ("create_time" not in record or _teaql_is_null(record["create_time"])): + record["create_time"] = Value.from_any(now) + context.record_fix_evidence("OrderSearchPreset", "create_time", "clock", "graphClock") + + if operation == "insert" and ("update_time" not in record or _teaql_is_null(record["update_time"])): + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("OrderSearchPreset", "update_time", "clock", "graphClock") + if operation == "insert" or operation == "update": + record["update_time"] = Value.from_any(now) + context.record_fix_evidence("OrderSearchPreset", "update_time", "clock", "graphClock") + + + if (operation == "insert" and "name" not in record) or ("name" in record and _teaql_is_null(record["name"])): + results.append(CheckResult("required", ObjectLocation().property("name"))) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 100)) + + if (operation == "insert" and "filter_json" not in record) or ("filter_json" in record and _teaql_is_null(record["filter_json"])): + results.append(CheckResult("required", ObjectLocation().property("filter_json"))) + if "filter_json" in record and _teaql_raw(record["filter_json"]) is not None and len(_teaql_raw(record["filter_json"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("filter_json"), _teaql_raw(record["filter_json"]), 100)) + + if (operation == "insert" and "request_id" not in record) or ("request_id" in record and _teaql_is_null(record["request_id"])): + results.append(CheckResult("required", ObjectLocation().property("request_id"))) + if "request_id" in record and _teaql_raw(record["request_id"]) is not None and len(_teaql_raw(record["request_id"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("request_id"), _teaql_raw(record["request_id"]), 100)) + + if (operation == "insert" and "owner_user_id" not in record) or ("owner_user_id" in record and _teaql_is_null(record["owner_user_id"])): + results.append(CheckResult("required", ObjectLocation().property("owner_user_id"))) + if "owner_user_id" in record and _teaql_raw(record["owner_user_id"]) is not None and len(_teaql_raw(record["owner_user_id"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("owner_user_id"), _teaql_raw(record["owner_user_id"]), 100)) + + if (operation == "insert" and "commerce_platform" not in record) or ("commerce_platform" in record and _teaql_is_null(record["commerce_platform"])): + results.append(CheckResult("required", ObjectLocation().property("commerce_platform"))) + + if (operation == "insert" and "create_time" not in record) or ("create_time" in record and _teaql_is_null(record["create_time"])): + results.append(CheckResult("required", ObjectLocation().property("create_time"))) + + if (operation == "insert" and "update_time" not in record) or ("update_time" in record and _teaql_is_null(record["update_time"])): + results.append(CheckResult("required", ObjectLocation().property("update_time"))) + + + +_CommercePlatform_DESCRIPTOR = (EntityDescriptor("CommercePlatform") + .table_name("commerce_platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("customer_list", "Customer").local("id").foreign("commerce_platform").many()).relation(RelationDescriptor("order_status_list", "OrderStatus").local("id").foreign("commerce_platform").many()).relation(RelationDescriptor("customer_order_list", "CustomerOrder").local("id").foreign("commerce_platform").many()).relation(RelationDescriptor("product_list", "Product").local("id").foreign("commerce_platform").many()).relation(RelationDescriptor("order_line_list", "OrderLine").local("id").foreign("commerce_platform").many()).relation(RelationDescriptor("order_search_preset_list", "OrderSearchPreset").local("id").foreign("commerce_platform").many()) +) + +_Customer_DESCRIPTOR = (EntityDescriptor("Customer") + .table_name("customer_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("email", DataType.Text).column_name("email").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")).relation(RelationDescriptor("customer_order_list", "CustomerOrder").local("id").foreign("customer").many()) +) + +_OrderStatus_DESCRIPTOR = (EntityDescriptor("OrderStatus") + .table_name("order_status_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("code", DataType.Text).column_name("code").required()).property(PropertyDescriptor("color", DataType.Text).column_name("color")).property(PropertyDescriptor("display_order", DataType.Decimal).column_name("display_order")).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")).relation(RelationDescriptor("customer_order_list", "CustomerOrder").local("id").foreign("status").many()) +) + +_CustomerOrder_DESCRIPTOR = (EntityDescriptor("CustomerOrder") + .table_name("customer_order_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("order_number", DataType.Text).column_name("order_number").required()).property(PropertyDescriptor("order_date", DataType.Date).column_name("order_date").required()).property(PropertyDescriptor("total_amount", DataType.Decimal).column_name("total_amount").required()).property(PropertyDescriptor("status", DataType.I64).column_name("status").required()).property(PropertyDescriptor("customer", DataType.I64).column_name("customer").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("status", "OrderStatus").local("status").foreign("id")).relation(RelationDescriptor("customer", "Customer").local("customer").foreign("id")).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")).relation(RelationDescriptor("order_line_list", "OrderLine").local("id").foreign("customer_order").many()) +) + +_Product_DESCRIPTOR = (EntityDescriptor("Product") + .table_name("product_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("sku", DataType.Text).column_name("sku").required()).property(PropertyDescriptor("image_url", DataType.Text).column_name("image_url")).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")).relation(RelationDescriptor("order_line_list", "OrderLine").local("id").foreign("product").many()) +) + +_OrderLine_DESCRIPTOR = (EntityDescriptor("OrderLine") + .table_name("order_line_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("customer_order", DataType.I64).column_name("customer_order").required()).property(PropertyDescriptor("product", DataType.I64).column_name("product").required()).property(PropertyDescriptor("product_name", DataType.Text).column_name("product_name").required()).property(PropertyDescriptor("sku", DataType.Text).column_name("sku").required()).property(PropertyDescriptor("quantity", DataType.I64).column_name("quantity").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("customer_order", "CustomerOrder").local("customer_order").foreign("id")).relation(RelationDescriptor("product", "Product").local("product").foreign("id")).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")) +) + +_OrderSearchPreset_DESCRIPTOR = (EntityDescriptor("OrderSearchPreset") + .table_name("order_search_preset_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("filter_json", DataType.Text).column_name("filter_json").required()).property(PropertyDescriptor("request_id", DataType.Text).column_name("request_id").required()).property(PropertyDescriptor("owner_user_id", DataType.Text).column_name("owner_user_id").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")) +) + +async def _ensure_generated_bootstrap_once(context): + previous_actor = context.user_identifier() if hasattr(context, 'user_identifier') else None + previous_category = context.get_resource('bootstrapCategory') + if hasattr(context, 'set_user_identifier'): + context.set_user_identifier('teaql-generated-bootstrap') + context.insert_resource('bootstrapCategory', 'runtime-bootstrap') + try: + commerce_platform_1 = await (Q.commerce_platforms().with_id_is(1).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if commerce_platform_1 is None: + commerce_platform_1 = CommercePlatform._teaql_new_with_fixed_id(1) + commerce_platform_1.update_name("Northwind Demo") + try: + await commerce_platform_1.audit_as('create model root CommercePlatform(1)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + commerce_platform_1 = await (Q.commerce_platforms().with_id_is(1).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if commerce_platform_1 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if commerce_platform_1 is None: + raise _teaql_create_error + context.with_active_root(ContextEntityRef("CommercePlatform", 1)) + order_status_1001 = await (Q.order_statuses().with_id_is(1001).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if order_status_1001 is None: + order_status_1001 = OrderStatus._teaql_new_with_fixed_id(1001) + order_status_1001.update_name("Pending") + order_status_1001.update_code("PENDING") + order_status_1001.update_color("#F59E0B") + order_status_1001.update_display_order(1) + order_status_1001.update_commerce_platform(CommercePlatform.refer(1)) + try: + await order_status_1001.audit_as('create model constant OrderStatus(1001)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + order_status_1001 = await (Q.order_statuses().with_id_is(1001).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if order_status_1001 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if order_status_1001 is None: + raise _teaql_create_error + _teaql_changed = False + if order_status_1001.name != "Pending": + order_status_1001.update_name("Pending") + _teaql_changed = True + if order_status_1001.code != "PENDING": + order_status_1001.update_code("PENDING") + _teaql_changed = True + if order_status_1001.color != "#F59E0B": + order_status_1001.update_color("#F59E0B") + _teaql_changed = True + if order_status_1001.displayOrder != 1: + order_status_1001.update_display_order(1) + _teaql_changed = True + if order_status_1001.commercePlatform != 1: + order_status_1001.update_commerce_platform(CommercePlatform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await order_status_1001.audit_as('reconcile model constant OrderStatus(1001)').save(context) + order_status_1002 = await (Q.order_statuses().with_id_is(1002).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if order_status_1002 is None: + order_status_1002 = OrderStatus._teaql_new_with_fixed_id(1002) + order_status_1002.update_name("Confirmed") + order_status_1002.update_code("CONFIRMED") + order_status_1002.update_color("#10B981") + order_status_1002.update_display_order(2) + order_status_1002.update_commerce_platform(CommercePlatform.refer(1)) + try: + await order_status_1002.audit_as('create model constant OrderStatus(1002)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + order_status_1002 = await (Q.order_statuses().with_id_is(1002).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if order_status_1002 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if order_status_1002 is None: + raise _teaql_create_error + _teaql_changed = False + if order_status_1002.name != "Confirmed": + order_status_1002.update_name("Confirmed") + _teaql_changed = True + if order_status_1002.code != "CONFIRMED": + order_status_1002.update_code("CONFIRMED") + _teaql_changed = True + if order_status_1002.color != "#10B981": + order_status_1002.update_color("#10B981") + _teaql_changed = True + if order_status_1002.displayOrder != 2: + order_status_1002.update_display_order(2) + _teaql_changed = True + if order_status_1002.commercePlatform != 1: + order_status_1002.update_commerce_platform(CommercePlatform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await order_status_1002.audit_as('reconcile model constant OrderStatus(1002)').save(context) + finally: + if hasattr(context, 'set_user_identifier'): + context.set_user_identifier(previous_actor) + context.insert_resource('bootstrapCategory', previous_category) + +async def _ensure_generated_bootstrap(context): + for _teaql_attempt in range(5): + try: + await _ensure_generated_bootstrap_once(context) + return + except Exception: + if _teaql_attempt == 4: + raise + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + + +# Passive generated manifest. Call ensure_schema() separately and explicitly. +GENERATED_RUNTIME_MODULE = (RuntimeModule().entity(CommercePlatform) + .schema_entity(_CommercePlatform_DESCRIPTOR) + .checker("CommercePlatform", _CommercePlatformChecker()) + .wire_metadata("CommercePlatform", create_wire_entity_metadata("CommercePlatform", ["id", "name", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})).entity(Customer) + .schema_entity(_Customer_DESCRIPTOR) + .checker("Customer", _CustomerChecker()) + .wire_metadata("Customer", create_wire_entity_metadata("Customer", ["id", "name", "email", "commerce_platform", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "email": ["email"], "commerce_platform": ["commerce_platform"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})).entity(OrderStatus) + .schema_entity(_OrderStatus_DESCRIPTOR) + .checker("OrderStatus", _OrderStatusChecker()) + .wire_metadata("OrderStatus", create_wire_entity_metadata("OrderStatus", ["id", "name", "code", "color", "display_order", "commerce_platform", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "code": ["code"], "color": ["color"], "display_order": ["display_order"], "commerce_platform": ["commerce_platform"], "version": ["version"]})).entity(CustomerOrder) + .schema_entity(_CustomerOrder_DESCRIPTOR) + .checker("CustomerOrder", _CustomerOrderChecker()) + .wire_metadata("CustomerOrder", create_wire_entity_metadata("CustomerOrder", ["id", "order_number", "order_date", "total_amount", "status", "customer", "commerce_platform", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "order_number": ["order_number"], "order_date": ["order_date"], "total_amount": ["total_amount"], "status": ["status"], "customer": ["customer"], "commerce_platform": ["commerce_platform"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})).entity(Product) + .schema_entity(_Product_DESCRIPTOR) + .checker("Product", _ProductChecker()) + .wire_metadata("Product", create_wire_entity_metadata("Product", ["id", "name", "sku", "image_url", "commerce_platform", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "sku": ["sku"], "image_url": ["image_url"], "commerce_platform": ["commerce_platform"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})).entity(OrderLine) + .schema_entity(_OrderLine_DESCRIPTOR) + .checker("OrderLine", _OrderLineChecker()) + .wire_metadata("OrderLine", create_wire_entity_metadata("OrderLine", ["id", "customer_order", "product", "product_name", "sku", "quantity", "commerce_platform", "create_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "customer_order": ["customer_order"], "product": ["product"], "product_name": ["product_name"], "sku": ["sku"], "quantity": ["quantity"], "commerce_platform": ["commerce_platform"], "create_time": ["create_time"], "version": ["version"]})).entity(OrderSearchPreset) + .schema_entity(_OrderSearchPreset_DESCRIPTOR) + .checker("OrderSearchPreset", _OrderSearchPresetChecker()) + .wire_metadata("OrderSearchPreset", create_wire_entity_metadata("OrderSearchPreset", ["id", "name", "filter_json", "request_id", "owner_user_id", "commerce_platform", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "filter_json": ["filter_json"], "request_id": ["request_id"], "owner_user_id": ["owner_user_id"], "commerce_platform": ["commerce_platform"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})) + .generated_bootstrap(_ensure_generated_bootstrap) +) \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql-i18n.json b/examples/order-management/python-lib-core/teaql-i18n.json new file mode 100644 index 0000000..f8c1a3b --- /dev/null +++ b/examples/order-management/python-lib-core/teaql-i18n.json @@ -0,0 +1,158 @@ +{ + "schema": "teaql.i18n/v1", + "defaultLocale": "en", + "locales": { + "de": { + "vocabulary": { + }, + "messages": { + } + }, + "ko": { + "vocabulary": { + }, + "messages": { + } + }, + "pt": { + "vocabulary": { + }, + "messages": { + } + }, + "zh-TW": { + "vocabulary": { + }, + "messages": { + } + }, + "fil": { + "vocabulary": { + }, + "messages": { + } + }, + "en": { + "vocabulary": { + "property.orderLine.quantity": "Quantity", + "property.customerOrder.createTime": "Create Time", + "property.customerOrder.updateTime": "Update Time", + "property.product.id": "Id", + "property.orderLine.createTime": "Create Time", + "entity.commercePlatform": "Commerce Platform", + "property.orderSearchPreset.requestId": "Request Id", + "property.product.version": "Version", + "property.product.createTime": "Create Time", + "property.orderSearchPreset.ownerUserId": "Owner User Id", + "entity.customerOrder": "Customer Order", + "entity.orderLine": "Order Line", + "property.product.name": "Name", + "property.orderLine.commercePlatform": "Commerce Platform", + "property.product.updateTime": "Update Time", + "property.customerOrder.totalAmount": "Total Amount", + "property.customerOrder.version": "Version", + "property.orderLine.productName": "Product Name", + "property.customer.version": "Version", + "property.orderSearchPreset.filterJson": "Filter Json", + "property.customer.id": "Id", + "property.orderLine.id": "Id", + "property.orderSearchPreset.name": "Name", + "property.customerOrder.customer": "Customer", + "property.orderLine.product": "Product", + "property.orderStatus.name": "Name", + "property.orderLine.sku": "Sku", + "property.customer.createTime": "Create Time", + "property.customerOrder.orderNumber": "Order Number", + "property.commercePlatform.name": "Name", + "property.orderStatus.commercePlatform": "Commerce Platform", + "property.product.commercePlatform": "Commerce Platform", + "property.commercePlatform.id": "Id", + "property.orderStatus.version": "Version", + "property.customer.updateTime": "Update Time", + "property.customer.commercePlatform": "Commerce Platform", + "property.orderStatus.code": "Code", + "property.orderSearchPreset.id": "Id", + "property.customerOrder.orderDate": "Order Date", + "property.commercePlatform.updateTime": "Update Time", + "property.customer.email": "Email", + "property.orderSearchPreset.createTime": "Create Time", + "property.orderStatus.displayOrder": "Display Order", + "property.orderSearchPreset.version": "Version", + "entity.customer": "Customer", + "property.customerOrder.id": "Id", + "entity.product": "Product", + "property.orderLine.customerOrder": "Customer Order", + "property.orderStatus.color": "Color", + "property.commercePlatform.createTime": "Create Time", + "property.customerOrder.commercePlatform": "Commerce Platform", + "property.orderSearchPreset.updateTime": "Update Time", + "property.product.sku": "Sku", + "property.orderSearchPreset.commercePlatform": "Commerce Platform", + "property.orderStatus.id": "Id", + "property.product.imageUrl": "Image Url", + "property.orderLine.version": "Version", + "entity.orderStatus": "Order Status", + "entity.orderSearchPreset": "Order Search Preset", + "property.customer.name": "Name", + "property.customerOrder.status": "Status", + "property.commercePlatform.version": "Version" + }, + "messages": { + } + }, + "fr": { + "vocabulary": { + }, + "messages": { + } + }, + "zh-CN": { + "vocabulary": { + }, + "messages": { + } + }, + "es": { + "vocabulary": { + }, + "messages": { + } + }, + "ar": { + "vocabulary": { + }, + "messages": { + } + }, + "vi": { + "vocabulary": { + }, + "messages": { + } + }, + "th": { + "vocabulary": { + }, + "messages": { + } + }, + "uk": { + "vocabulary": { + }, + "messages": { + } + }, + "ja": { + "vocabulary": { + }, + "messages": { + } + }, + "id": { + "vocabulary": { + }, + "messages": { + } + } + } +} \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql/__init__.py b/examples/order-management/python-lib-core/teaql/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/order-management/python-lib-core/teaql/core/__init__.py b/examples/order-management/python-lib-core/teaql/core/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/order-management/python-lib-core/teaql/core/expr.py b/examples/order-management/python-lib-core/teaql/core/expr.py deleted file mode 100644 index b77f00c..0000000 --- a/examples/order-management/python-lib-core/teaql/core/expr.py +++ /dev/null @@ -1,764 +0,0 @@ -import copy -import json -import os -import re -import tempfile -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse - -ENTITY_SCHEMAS = { -"CommercePlatform": { - "table": "commerce_platform_data", - "columns": {"id": "integer", "name": "text", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_list": {"target_entity": "Customer", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_status_list": {"target_entity": "OrderStatus", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "product_list": {"target_entity": "Product", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_search_preset_list": {"target_entity": "OrderSearchPreset", "local_key": "id", "foreign_key": "commerce_platform", "many": True}}, -}, -"Customer": { - "table": "customer_data", - "columns": {"id": "integer", "name": "text", "email": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "customer", "many": True}}, -}, -"OrderStatus": { - "table": "order_status_data", - "columns": {"id": "integer", "name": "text", "code": "text", "color": "text", "display_order": "integer", "commerce_platform": "integer", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "status", "many": True}}, -}, -"CustomerOrder": { - "table": "customer_order_data", - "columns": {"id": "integer", "order_number": "text", "order_date": "date", "total_amount": "integer", "status": "integer", "customer": "integer", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "customer_order", "many": True}}, -}, -"Product": { - "table": "product_data", - "columns": {"id": "integer", "name": "text", "sku": "text", "image_url": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "product", "many": True}}, -}, -"OrderLine": { - "table": "order_line_data", - "columns": {"id": "integer", "customer_order": "integer", "product": "integer", "product_name": "text", "sku": "text", "quantity": "integer", "commerce_platform": "integer", "create_time": "date", "version": "integer"}, - "relations": {}, -}, -"OrderSearchPreset": { - "table": "order_search_preset_data", - "columns": {"id": "integer", "name": "text", "filter_json": "text", "request_id": "text", "owner_user_id": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._relations = [] - self._partition_by = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): self._limit = n - def offset(self, n): self._offset = n - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - -class QueryRequest: - def __init__(self, query): - self.query = query - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._load() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - del table[str(record_id)] - self._persist() - result = {"success": True, "id": record_id, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query = req.query - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - return type('QueryResult', (object,), {'rows': rows[start:end]}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - ) - return table - - async def ensure_schema(self): - connection = await self._connect() - try: - async with connection.transaction(): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - finally: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "entity VARCHAR(255) PRIMARY KEY, next_id BIGINT NOT NULL)" - ) - if self.database_kind == "postgres": - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES ($1, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - if self.database_kind == "mysql": - await connection.execute( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (%s, 1000) " - "ON DUPLICATE KEY UPDATE next_id = LAST_INSERT_ID(next_id + 1)", - entity, - ) - return await connection.fetch_value( - "SELECT next_id FROM teaql_id_space WHERE entity = %s", entity - ) - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (?, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - - async def mutate(self, context, req): - command = req.cmd - connection = await self._connect() - try: - async with connection.transaction(): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})", - *params, - ) - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - row = await connection.fetch_one( - f"SELECT {version} FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = {"success": True, "id": command.pk, "version": row["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"DELETE FROM {self._identifier(table)} WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - result = {"success": True, "id": command.pk, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - async def query(self, context, req): - query = req.query - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator == "in": - values = list(expression.get("value") or []) - if not values: - predicates.append("1 = 0") - continue - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - predicates.append(f"{field} IN ({', '.join(placeholders)})") - continue - params.append(self._normalize(expression.get("value"))) - placeholder = self._placeholder(len(params)) - if operator == "eq": - predicates.append(f"{field} = {placeholder}") - elif operator == "contain": - predicates.append(self._contains_predicate(field, placeholder)) - elif operator == "gte": - predicates.append(f"{field} >= {placeholder}") - elif operator == "lte": - predicates.append(f"{field} <= {placeholder}") - else: - raise ValueError(f"Unsupported filter operator: {operator}") - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - return type('QueryResult', (object,), {'rows': rows}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query.entity = relation["target_entity"] - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql/core/mutation.py b/examples/order-management/python-lib-core/teaql/core/mutation.py deleted file mode 100644 index b77f00c..0000000 --- a/examples/order-management/python-lib-core/teaql/core/mutation.py +++ /dev/null @@ -1,764 +0,0 @@ -import copy -import json -import os -import re -import tempfile -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse - -ENTITY_SCHEMAS = { -"CommercePlatform": { - "table": "commerce_platform_data", - "columns": {"id": "integer", "name": "text", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_list": {"target_entity": "Customer", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_status_list": {"target_entity": "OrderStatus", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "product_list": {"target_entity": "Product", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_search_preset_list": {"target_entity": "OrderSearchPreset", "local_key": "id", "foreign_key": "commerce_platform", "many": True}}, -}, -"Customer": { - "table": "customer_data", - "columns": {"id": "integer", "name": "text", "email": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "customer", "many": True}}, -}, -"OrderStatus": { - "table": "order_status_data", - "columns": {"id": "integer", "name": "text", "code": "text", "color": "text", "display_order": "integer", "commerce_platform": "integer", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "status", "many": True}}, -}, -"CustomerOrder": { - "table": "customer_order_data", - "columns": {"id": "integer", "order_number": "text", "order_date": "date", "total_amount": "integer", "status": "integer", "customer": "integer", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "customer_order", "many": True}}, -}, -"Product": { - "table": "product_data", - "columns": {"id": "integer", "name": "text", "sku": "text", "image_url": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "product", "many": True}}, -}, -"OrderLine": { - "table": "order_line_data", - "columns": {"id": "integer", "customer_order": "integer", "product": "integer", "product_name": "text", "sku": "text", "quantity": "integer", "commerce_platform": "integer", "create_time": "date", "version": "integer"}, - "relations": {}, -}, -"OrderSearchPreset": { - "table": "order_search_preset_data", - "columns": {"id": "integer", "name": "text", "filter_json": "text", "request_id": "text", "owner_user_id": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._relations = [] - self._partition_by = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): self._limit = n - def offset(self, n): self._offset = n - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - -class QueryRequest: - def __init__(self, query): - self.query = query - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._load() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - del table[str(record_id)] - self._persist() - result = {"success": True, "id": record_id, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query = req.query - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - return type('QueryResult', (object,), {'rows': rows[start:end]}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - ) - return table - - async def ensure_schema(self): - connection = await self._connect() - try: - async with connection.transaction(): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - finally: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "entity VARCHAR(255) PRIMARY KEY, next_id BIGINT NOT NULL)" - ) - if self.database_kind == "postgres": - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES ($1, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - if self.database_kind == "mysql": - await connection.execute( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (%s, 1000) " - "ON DUPLICATE KEY UPDATE next_id = LAST_INSERT_ID(next_id + 1)", - entity, - ) - return await connection.fetch_value( - "SELECT next_id FROM teaql_id_space WHERE entity = %s", entity - ) - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (?, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - - async def mutate(self, context, req): - command = req.cmd - connection = await self._connect() - try: - async with connection.transaction(): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})", - *params, - ) - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - row = await connection.fetch_one( - f"SELECT {version} FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = {"success": True, "id": command.pk, "version": row["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"DELETE FROM {self._identifier(table)} WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - result = {"success": True, "id": command.pk, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - async def query(self, context, req): - query = req.query - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator == "in": - values = list(expression.get("value") or []) - if not values: - predicates.append("1 = 0") - continue - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - predicates.append(f"{field} IN ({', '.join(placeholders)})") - continue - params.append(self._normalize(expression.get("value"))) - placeholder = self._placeholder(len(params)) - if operator == "eq": - predicates.append(f"{field} = {placeholder}") - elif operator == "contain": - predicates.append(self._contains_predicate(field, placeholder)) - elif operator == "gte": - predicates.append(f"{field} >= {placeholder}") - elif operator == "lte": - predicates.append(f"{field} <= {placeholder}") - else: - raise ValueError(f"Unsupported filter operator: {operator}") - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - return type('QueryResult', (object,), {'rows': rows}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query.entity = relation["target_entity"] - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql/core/query.py b/examples/order-management/python-lib-core/teaql/core/query.py deleted file mode 100644 index b77f00c..0000000 --- a/examples/order-management/python-lib-core/teaql/core/query.py +++ /dev/null @@ -1,764 +0,0 @@ -import copy -import json -import os -import re -import tempfile -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse - -ENTITY_SCHEMAS = { -"CommercePlatform": { - "table": "commerce_platform_data", - "columns": {"id": "integer", "name": "text", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_list": {"target_entity": "Customer", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_status_list": {"target_entity": "OrderStatus", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "product_list": {"target_entity": "Product", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_search_preset_list": {"target_entity": "OrderSearchPreset", "local_key": "id", "foreign_key": "commerce_platform", "many": True}}, -}, -"Customer": { - "table": "customer_data", - "columns": {"id": "integer", "name": "text", "email": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "customer", "many": True}}, -}, -"OrderStatus": { - "table": "order_status_data", - "columns": {"id": "integer", "name": "text", "code": "text", "color": "text", "display_order": "integer", "commerce_platform": "integer", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "status", "many": True}}, -}, -"CustomerOrder": { - "table": "customer_order_data", - "columns": {"id": "integer", "order_number": "text", "order_date": "date", "total_amount": "integer", "status": "integer", "customer": "integer", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "customer_order", "many": True}}, -}, -"Product": { - "table": "product_data", - "columns": {"id": "integer", "name": "text", "sku": "text", "image_url": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "product", "many": True}}, -}, -"OrderLine": { - "table": "order_line_data", - "columns": {"id": "integer", "customer_order": "integer", "product": "integer", "product_name": "text", "sku": "text", "quantity": "integer", "commerce_platform": "integer", "create_time": "date", "version": "integer"}, - "relations": {}, -}, -"OrderSearchPreset": { - "table": "order_search_preset_data", - "columns": {"id": "integer", "name": "text", "filter_json": "text", "request_id": "text", "owner_user_id": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._relations = [] - self._partition_by = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): self._limit = n - def offset(self, n): self._offset = n - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - -class QueryRequest: - def __init__(self, query): - self.query = query - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._load() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - del table[str(record_id)] - self._persist() - result = {"success": True, "id": record_id, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query = req.query - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - return type('QueryResult', (object,), {'rows': rows[start:end]}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - ) - return table - - async def ensure_schema(self): - connection = await self._connect() - try: - async with connection.transaction(): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - finally: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "entity VARCHAR(255) PRIMARY KEY, next_id BIGINT NOT NULL)" - ) - if self.database_kind == "postgres": - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES ($1, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - if self.database_kind == "mysql": - await connection.execute( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (%s, 1000) " - "ON DUPLICATE KEY UPDATE next_id = LAST_INSERT_ID(next_id + 1)", - entity, - ) - return await connection.fetch_value( - "SELECT next_id FROM teaql_id_space WHERE entity = %s", entity - ) - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (?, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - - async def mutate(self, context, req): - command = req.cmd - connection = await self._connect() - try: - async with connection.transaction(): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})", - *params, - ) - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - row = await connection.fetch_one( - f"SELECT {version} FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = {"success": True, "id": command.pk, "version": row["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"DELETE FROM {self._identifier(table)} WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - result = {"success": True, "id": command.pk, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - async def query(self, context, req): - query = req.query - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator == "in": - values = list(expression.get("value") or []) - if not values: - predicates.append("1 = 0") - continue - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - predicates.append(f"{field} IN ({', '.join(placeholders)})") - continue - params.append(self._normalize(expression.get("value"))) - placeholder = self._placeholder(len(params)) - if operator == "eq": - predicates.append(f"{field} = {placeholder}") - elif operator == "contain": - predicates.append(self._contains_predicate(field, placeholder)) - elif operator == "gte": - predicates.append(f"{field} >= {placeholder}") - elif operator == "lte": - predicates.append(f"{field} <= {placeholder}") - else: - raise ValueError(f"Unsupported filter operator: {operator}") - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - return type('QueryResult', (object,), {'rows': rows}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query.entity = relation["target_entity"] - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql/core/value.py b/examples/order-management/python-lib-core/teaql/core/value.py deleted file mode 100644 index b77f00c..0000000 --- a/examples/order-management/python-lib-core/teaql/core/value.py +++ /dev/null @@ -1,764 +0,0 @@ -import copy -import json -import os -import re -import tempfile -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse - -ENTITY_SCHEMAS = { -"CommercePlatform": { - "table": "commerce_platform_data", - "columns": {"id": "integer", "name": "text", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_list": {"target_entity": "Customer", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_status_list": {"target_entity": "OrderStatus", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "product_list": {"target_entity": "Product", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_search_preset_list": {"target_entity": "OrderSearchPreset", "local_key": "id", "foreign_key": "commerce_platform", "many": True}}, -}, -"Customer": { - "table": "customer_data", - "columns": {"id": "integer", "name": "text", "email": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "customer", "many": True}}, -}, -"OrderStatus": { - "table": "order_status_data", - "columns": {"id": "integer", "name": "text", "code": "text", "color": "text", "display_order": "integer", "commerce_platform": "integer", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "status", "many": True}}, -}, -"CustomerOrder": { - "table": "customer_order_data", - "columns": {"id": "integer", "order_number": "text", "order_date": "date", "total_amount": "integer", "status": "integer", "customer": "integer", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "customer_order", "many": True}}, -}, -"Product": { - "table": "product_data", - "columns": {"id": "integer", "name": "text", "sku": "text", "image_url": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "product", "many": True}}, -}, -"OrderLine": { - "table": "order_line_data", - "columns": {"id": "integer", "customer_order": "integer", "product": "integer", "product_name": "text", "sku": "text", "quantity": "integer", "commerce_platform": "integer", "create_time": "date", "version": "integer"}, - "relations": {}, -}, -"OrderSearchPreset": { - "table": "order_search_preset_data", - "columns": {"id": "integer", "name": "text", "filter_json": "text", "request_id": "text", "owner_user_id": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._relations = [] - self._partition_by = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): self._limit = n - def offset(self, n): self._offset = n - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - -class QueryRequest: - def __init__(self, query): - self.query = query - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._load() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - del table[str(record_id)] - self._persist() - result = {"success": True, "id": record_id, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query = req.query - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - return type('QueryResult', (object,), {'rows': rows[start:end]}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - ) - return table - - async def ensure_schema(self): - connection = await self._connect() - try: - async with connection.transaction(): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - finally: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "entity VARCHAR(255) PRIMARY KEY, next_id BIGINT NOT NULL)" - ) - if self.database_kind == "postgres": - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES ($1, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - if self.database_kind == "mysql": - await connection.execute( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (%s, 1000) " - "ON DUPLICATE KEY UPDATE next_id = LAST_INSERT_ID(next_id + 1)", - entity, - ) - return await connection.fetch_value( - "SELECT next_id FROM teaql_id_space WHERE entity = %s", entity - ) - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (?, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - - async def mutate(self, context, req): - command = req.cmd - connection = await self._connect() - try: - async with connection.transaction(): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})", - *params, - ) - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - row = await connection.fetch_one( - f"SELECT {version} FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = {"success": True, "id": command.pk, "version": row["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"DELETE FROM {self._identifier(table)} WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - result = {"success": True, "id": command.pk, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - async def query(self, context, req): - query = req.query - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator == "in": - values = list(expression.get("value") or []) - if not values: - predicates.append("1 = 0") - continue - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - predicates.append(f"{field} IN ({', '.join(placeholders)})") - continue - params.append(self._normalize(expression.get("value"))) - placeholder = self._placeholder(len(params)) - if operator == "eq": - predicates.append(f"{field} = {placeholder}") - elif operator == "contain": - predicates.append(self._contains_predicate(field, placeholder)) - elif operator == "gte": - predicates.append(f"{field} >= {placeholder}") - elif operator == "lte": - predicates.append(f"{field} <= {placeholder}") - else: - raise ValueError(f"Unsupported filter operator: {operator}") - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - return type('QueryResult', (object,), {'rows': rows}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query.entity = relation["target_entity"] - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/order-management/python-lib-core/teaql/data_service.py b/examples/order-management/python-lib-core/teaql/data_service.py deleted file mode 100644 index a63be3a..0000000 --- a/examples/order-management/python-lib-core/teaql/data_service.py +++ /dev/null @@ -1,767 +0,0 @@ -import copy -import json -import os -import re -import tempfile -from teaql.runtime import _SCHEMA_INVOCATION -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse - -ENTITY_SCHEMAS = { -"CommercePlatform": { - "table": "commerce_platform_data", - "columns": {"id": "integer", "name": "text", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_list": {"target_entity": "Customer", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_status_list": {"target_entity": "OrderStatus", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "product_list": {"target_entity": "Product", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "commerce_platform", "many": True}, "order_search_preset_list": {"target_entity": "OrderSearchPreset", "local_key": "id", "foreign_key": "commerce_platform", "many": True}}, -}, -"Customer": { - "table": "customer_data", - "columns": {"id": "integer", "name": "text", "email": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "customer", "many": True}}, -}, -"OrderStatus": { - "table": "order_status_data", - "columns": {"id": "integer", "name": "text", "code": "text", "color": "text", "display_order": "integer", "commerce_platform": "integer", "version": "integer"}, - "relations": {"customer_order_list": {"target_entity": "CustomerOrder", "local_key": "id", "foreign_key": "status", "many": True}}, -}, -"CustomerOrder": { - "table": "customer_order_data", - "columns": {"id": "integer", "order_number": "text", "order_date": "date", "total_amount": "integer", "status": "integer", "customer": "integer", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "customer_order", "many": True}}, -}, -"Product": { - "table": "product_data", - "columns": {"id": "integer", "name": "text", "sku": "text", "image_url": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {"order_line_list": {"target_entity": "OrderLine", "local_key": "id", "foreign_key": "product", "many": True}}, -}, -"OrderLine": { - "table": "order_line_data", - "columns": {"id": "integer", "customer_order": "integer", "product": "integer", "product_name": "text", "sku": "text", "quantity": "integer", "commerce_platform": "integer", "create_time": "date", "version": "integer"}, - "relations": {}, -}, -"OrderSearchPreset": { - "table": "order_search_preset_data", - "columns": {"id": "integer", "name": "text", "filter_json": "text", "request_id": "text", "owner_user_id": "text", "commerce_platform": "integer", "create_time": "date", "update_time": "date", "version": "integer"}, - "relations": {}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._relations = [] - self._partition_by = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): self._limit = n - def offset(self, n): self._offset = n - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - -class QueryRequest: - def __init__(self, query): - self.query = query - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._load() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - self._persist() - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - del table[str(record_id)] - self._persist() - result = {"success": True, "id": record_id, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query = req.query - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - return type('QueryResult', (object,), {'rows': rows[start:end]}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - connection = await self._connect() - try: - async with connection.transaction(): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - finally: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "entity VARCHAR(255) PRIMARY KEY, next_id BIGINT NOT NULL)" - ) - if self.database_kind == "postgres": - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES ($1, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - if self.database_kind == "mysql": - await connection.execute( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (%s, 1000) " - "ON DUPLICATE KEY UPDATE next_id = LAST_INSERT_ID(next_id + 1)", - entity, - ) - return await connection.fetch_value( - "SELECT next_id FROM teaql_id_space WHERE entity = %s", entity - ) - return await connection.fetch_value( - "INSERT INTO teaql_id_space(entity, next_id) VALUES (?, 1000) " - "ON CONFLICT(entity) DO UPDATE SET next_id = teaql_id_space.next_id + 1 " - "RETURNING next_id", - entity, - ) - - async def mutate(self, context, req): - command = req.cmd - connection = await self._connect() - try: - async with connection.transaction(): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})", - *params, - ) - result = {"success": True, "id": record_id, "version": record["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - row = await connection.fetch_one( - f"SELECT {version} FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = {"success": True, "id": command.pk, "version": row["version"]} - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - affected = await connection.execute( - f"DELETE FROM {self._identifier(table)} WHERE {' AND '.join(predicates)}", - *params, - ) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - result = {"success": True, "id": command.pk, "deleted": True} - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - async def query(self, context, req): - query = req.query - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator == "in": - values = list(expression.get("value") or []) - if not values: - predicates.append("1 = 0") - continue - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - predicates.append(f"{field} IN ({', '.join(placeholders)})") - continue - params.append(self._normalize(expression.get("value"))) - placeholder = self._placeholder(len(params)) - if operator == "eq": - predicates.append(f"{field} = {placeholder}") - elif operator == "contain": - predicates.append(self._contains_predicate(field, placeholder)) - elif operator == "gte": - predicates.append(f"{field} >= {placeholder}") - elif operator == "lte": - predicates.append(f"{field} <= {placeholder}") - else: - raise ValueError(f"Unsupported filter operator: {operator}") - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - return type('QueryResult', (object,), {'rows': rows}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query.entity = relation["target_entity"] - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) diff --git a/examples/order-management/python-lib-core/teaql/runtime.py b/examples/order-management/python-lib-core/teaql/runtime.py deleted file mode 100644 index 4b9eff1..0000000 --- a/examples/order-management/python-lib-core/teaql/runtime.py +++ /dev/null @@ -1,82 +0,0 @@ -from dataclasses import dataclass - -_SCHEMA_INVOCATION = object() - -@dataclass(frozen=True) -class RawAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - changes: tuple - -@dataclass(frozen=True) -class SafeAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - fields: tuple - -class UserContext: - """Runtime dependencies and trusted request state initialized by the server.""" - - def __init__(self): - self._resources = {} - self._standard_audit_sink = None - self._app_audit_sink = None - self._audit_policies = {} - - @classmethod - def new(cls): - return cls() - - def insert_resource(self, resource_type, resource): - self._resources[resource_type] = resource - return self - - def get_resource(self, resource_type): - return self._resources.get(resource_type) - - def require_resource(self, resource_type): - resource = self.get_resource(resource_type) - if resource is None: - raise RuntimeError(f"Required UserContext resource is missing: {resource_type}") - return resource - - async def ensure_schema(self): - """Reconcile schema through this context's configured data service.""" - await self.require_resource("dataService")._ensure_schema(self, _SCHEMA_INVOCATION) - - def initialize_audit(self, standard_sink, app_sink=None): - self._standard_audit_sink = standard_sink - self._app_audit_sink = app_sink - return self - - def configure_audit_policy(self, entity, mask_fields=(), max_length=None): - self._audit_policies[entity] = (frozenset(mask_fields), max_length) - return self - - async def emit_mutation_audit(self, req, result): - command = req.cmd - values = getattr(command, "payload", getattr(command, "values", {})) - kind = "created" if hasattr(command, "payload") else "updated" if hasattr(command, "values") else "deleted" - raw = RawAuditEvent(kind, command.entity, result.get("id"), req.comment, - tuple((name, None, value) for name, value in values.items())) - if self._standard_audit_sink is not None: - emitted = self._standard_audit_sink.on_event(self, raw) - if hasattr(emitted, "__await__"): await emitted - if self._app_audit_sink is not None: - masks, limit = self._audit_policies.get(command.entity, (frozenset(), None)) - fields = [] - for name, _, raw_value in raw.changes: - value = None if raw_value is None else str(raw_value) - masked = name in masks - if value is not None and masked: - value = "*" * len(value) if len(value) < 8 else value[:2] + "*" * (len(value) - 4) + value[-2:] - truncated = value is not None and limit is not None and len(value) > limit - if truncated: value = "*" * limit if limit <= 3 else value[:limit - 3] + "..." - fields.append((name, value, masked, truncated)) - safe = SafeAuditEvent(kind, command.entity, result.get("id"), req.comment, tuple(fields)) - emitted = self._app_audit_sink.on_safe_event(self, safe) - if hasattr(emitted, "__await__"): await emitted diff --git a/examples/school-management/pyproject.toml b/examples/school-management/pyproject.toml index 6a52cc1..a7537f1 100644 --- a/examples/school-management/pyproject.toml +++ b/examples/school-management/pyproject.toml @@ -2,15 +2,15 @@ name = "school-management-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] [tool.setuptools.packages.find] where = ["."] -include = ["models*", "requests*", "teaql*"] +include = ["models*", "requests*"] [build-system] requires = ["setuptools>=42"] -build-backend = "setuptools.build_meta" +build-backend = "setuptools.build_meta" \ No newline at end of file diff --git a/examples/school-management/runtime_module.py b/examples/school-management/runtime_module.py index dc2aca6..cb468ae 100644 --- a/examples/school-management/runtime_module.py +++ b/examples/school-management/runtime_module.py @@ -1,6 +1,8 @@ import asyncio from datetime import datetime, timezone -from teaql.runtime import CheckResult, ContextEntityRef, ObjectLocation, RuntimeModule +from teaql.runtime import CheckResult, ContextEntityRef, JsonFieldNamingProfile, ObjectLocation, RuntimeModule, create_wire_entity_metadata +from teaql.core.meta import EntityDescriptor, PropertyDescriptor, RelationDescriptor +from teaql.core.value import DataType from Q import Q from teaql.core.value import Value try: @@ -136,6 +138,18 @@ def check_and_fix(self, context, record, location, results): +_Platform_DESCRIPTOR = (EntityDescriptor("Platform") + .table_name("platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("base_url", DataType.Text).column_name("base_url").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("school_type_list", "SchoolType").local("id").foreign("platform").many()).relation(RelationDescriptor("school_list", "School").local("id").foreign("platform").many()) +) + +_SchoolType_DESCRIPTOR = (EntityDescriptor("SchoolType") + .table_name("school_type_data").property(PropertyDescriptor("platform", DataType.I64).column_name("platform").required()).property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("code", DataType.Text).column_name("code").required()).property(PropertyDescriptor("display_order", DataType.Decimal).column_name("display_order").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")).relation(RelationDescriptor("school_list", "School").local("id").foreign("school_type").many()) +) + +_School_DESCRIPTOR = (EntityDescriptor("School") + .table_name("school_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").is_id().required()).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").required()).property(PropertyDescriptor("school_type", DataType.I64).column_name("school_type").required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").required()).property(PropertyDescriptor("address", DataType.Text).column_name("address").required()).property(PropertyDescriptor("established_date", DataType.Date).column_name("established_date").required()).property(PropertyDescriptor("student_capacity", DataType.I64).column_name("student_capacity").required()).property(PropertyDescriptor("active", DataType.Bool).column_name("active").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").is_version().required()).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")).relation(RelationDescriptor("school_type", "SchoolType").local("school_type").foreign("id")) +) + async def _ensure_generated_bootstrap_once(context): previous_actor = context.user_identifier() if hasattr(context, 'user_identifier') else None previous_category = context.get_resource('bootstrapCategory') @@ -244,8 +258,14 @@ async def _ensure_generated_bootstrap(context): # Passive generated manifest. Call ensure_schema() separately and explicitly. GENERATED_RUNTIME_MODULE = (RuntimeModule().entity(Platform) - .checker("Platform", _PlatformChecker()).entity(SchoolType) - .checker("SchoolType", _SchoolTypeChecker()).entity(School) + .schema_entity(_Platform_DESCRIPTOR) + .checker("Platform", _PlatformChecker()) + .wire_metadata("Platform", create_wire_entity_metadata("Platform", ["id", "name", "base_url", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "base_url": ["base_url"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})).entity(SchoolType) + .schema_entity(_SchoolType_DESCRIPTOR) + .checker("SchoolType", _SchoolTypeChecker()) + .wire_metadata("SchoolType", create_wire_entity_metadata("SchoolType", ["platform", "id", "name", "code", "display_order", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"platform": ["platform"], "id": ["id"], "name": ["name"], "code": ["code"], "display_order": ["display_order"], "version": ["version"]})).entity(School) + .schema_entity(_School_DESCRIPTOR) .checker("School", _SchoolChecker()) + .wire_metadata("School", create_wire_entity_metadata("School", ["id", "platform", "school_type", "name", "address", "established_date", "student_capacity", "active", "create_time", "update_time", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "platform": ["platform"], "school_type": ["school_type"], "name": ["name"], "address": ["address"], "established_date": ["established_date"], "student_capacity": ["student_capacity"], "active": ["active"], "create_time": ["create_time"], "update_time": ["update_time"], "version": ["version"]})) .generated_bootstrap(_ensure_generated_bootstrap) ) \ No newline at end of file diff --git a/examples/school-management/teaql-i18n.json b/examples/school-management/teaql-i18n.json index 474f7b0..21b3f9b 100644 --- a/examples/school-management/teaql-i18n.json +++ b/examples/school-management/teaql-i18n.json @@ -72,7 +72,6 @@ }, "zh-CN": { "vocabulary": { - "property.platform.baseUrl": "https" }, "messages": { } diff --git a/examples/school-management/teaql/__init__.py b/examples/school-management/teaql/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/school-management/teaql/core/__init__.py b/examples/school-management/teaql/core/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/examples/school-management/teaql/core/expr.py b/examples/school-management/teaql/core/expr.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/core/expr.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/core/list.py b/examples/school-management/teaql/core/list.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/core/list.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/core/mutation.py b/examples/school-management/teaql/core/mutation.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/core/mutation.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/core/query.py b/examples/school-management/teaql/core/query.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/core/query.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/core/value.py b/examples/school-management/teaql/core/value.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/core/value.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/data_service.py b/examples/school-management/teaql/data_service.py deleted file mode 100644 index 0660d9f..0000000 --- a/examples/school-management/teaql/data_service.py +++ /dev/null @@ -1,1381 +0,0 @@ -import copy -import json -import os -import re -import tempfile -import hashlib -import time -import asyncio -from datetime import date, datetime -from decimal import Decimal -from urllib.parse import parse_qs, unquote, urlparse -from dataclasses import dataclass -from typing import Any, Callable, Dict, Generic, Iterable, Optional, TypeVar -from teaql.runtime import SqlLogOperation, _SCHEMA_INVOCATION - -TPage = TypeVar("TPage") - -class SmartList(list[TPage], Generic[TPage]): - def __init__(self, data: Iterable[TPage] = (), facets: Optional[Dict[str, Any]] = None, - total_count: Optional[int] = None): - super().__init__(data) - self.facets = facets or {} - self.total_count = len(self) if total_count is None else total_count - - @property - def data(self) -> "SmartList[TPage]": - return self - - def facet(self, name: str) -> Any: - return self.facets.get(name) - - def map(self, mapper: Callable[[TPage], Any]) -> "SmartList[Any]": - return SmartList((mapper(item) for item in self), self.facets, self.total_count) - - def filter(self, predicate: Callable[[TPage], bool]) -> "SmartList[TPage]": - return SmartList((item for item in self if predicate(item)), self.facets, self.total_count) - - def first(self) -> Optional[TPage]: - return self[0] if self else None - - def last(self) -> Optional[TPage]: - return self[-1] if self else None - -@dataclass(frozen=True) -class TeaQLPage(Generic[TPage]): - data: SmartList[TPage] - total_count: int - offset: int - limit: int - -ENTITY_SCHEMAS = { -"Platform": { - "table": "platform_data", - "columns": {"id": "integer", "name": "text", "base_url": "text", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "name": True, "base_url": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{}, **{"school_type_list": {"target_entity": "SchoolType", "local_key": "id", "foreign_key": "platform", "many": True}, "school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "platform", "many": True}}}, -}, -"SchoolType": { - "table": "school_type_data", - "columns": {"platform": "integer", "id": "integer", "name": "text", "code": "text", "display_order": "decimal", "version": "integer"}, - "required": {"platform": True, "id": True, "name": True, "code": True, "display_order": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}}, **{"school_list": {"target_entity": "School", "local_key": "id", "foreign_key": "school_type", "many": True}}}, -}, -"School": { - "table": "school_data", - "columns": {"id": "integer", "platform": "integer", "school_type": "integer", "name": "text", "address": "text", "established_date": "date", "student_capacity": "integer", "active": "bool", "create_time": "datetime", "update_time": "datetime", "version": "integer"}, - "required": {"id": True, "platform": True, "school_type": True, "name": True, "address": True, "established_date": True, "student_capacity": True, "active": True, "create_time": True, "update_time": True, "version": True}, - "relations": {**{"platform": {"target_entity": "Platform", "local_key": "platform", "foreign_key": "id", "many": False}, "school_type": {"target_entity": "SchoolType", "local_key": "school_type", "foreign_key": "id", "many": False}}, **{}}, -} -} - -class Value: - @staticmethod - def Text(val): return val - @staticmethod - def I64(val): return val - @staticmethod - def F64(val): return val - @staticmethod - def Decimal(val): return val - @staticmethod - def Date(val): return val - @staticmethod - def DateTime(val): return val - @staticmethod - def Bool(val): return val - @staticmethod - def JSON(val): return val - @staticmethod - def Object(val): return val - @staticmethod - def from_any(val): return val - -class SelectQuery: - def __init__(self, entity): - self.entity = entity - self._comment = None - self._purpose = None - self._trace_path = [] - self._limit = None - self._offset = None - self._order_by = [] - self._group_by = [] - self._aggregates = [] - self._filters = [] - self._projection = [] - self._relations = [] - self._relation_aggregates = [] - self._facets = [] - self._partition_by = None - self._top_n_probe_parent_threshold = None - self._continuous_page_fetch_options = None - self.id_set_pagination = None - - def comment(self, c): self._comment = c - def purpose(self, p): self._purpose = p - def limit(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 1: - raise ValueError("QUERY_INVALID_LIMIT: limit must be a positive integer") - if n > 10_000: raise ValueError("QUERY_HARD_LIMIT_EXCEEDED: limit exceeds 10000") - self._limit = n - return self - def offset(self, n): - if not isinstance(n, int) or isinstance(n, bool) or n < 0: - raise ValueError("QUERY_INVALID_OFFSET: offset must be a non-negative integer") - self._offset = n - return self - def order_by(self, f, d): self._order_by.append((f, d)) - def group_by(self, f): self._group_by.append(f) - def count_field(self, f, n): self._aggregates.append(("count", f, n)) - def aggregate(self, func, field, ret_name): self._aggregates.append((func, field, ret_name)) - def and_filter(self, expr): self._filters.append(expr) - def project(self, *fields): - for field in fields: - if field not in self._projection: self._projection.append(field) - return self - def relation_query(self, name, query): self._relations.append({"name": name, "query": query}) - def top_n_probe_parent_threshold(self, threshold): - if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0: - raise ValueError("Top-N probe parent threshold must not be negative") - self._top_n_probe_parent_threshold = threshold - return self - def relation_aggregate(self, relation_name, alias, query, single_result=True): - self._relation_aggregates.append({ - "relation_name": relation_name, "alias": alias, - "query": query, "single_result": single_result}) - return self - def facet_by(self, name, relation_name, query, include_all_facets=True): - self._facets.append({ - "name": name, "relation_name": relation_name, "query": query, - "include_all_facets": include_all_facets}) - return self - def for_exact_count(self, alias="__teaql_total"): - query = copy.deepcopy(self) - query._projection = [] - query._relations = [] - query._facets = [] - query._order_by = [] - query._offset = None - query._limit = None - query._group_by = [] - query._aggregates = [("count", "id", alias)] - return query - def optimize_for_continuous_page_fetch(self): - return self.optimize_for_continuous_page_fetch_with("default", 600) - def optimize_for_continuous_page_fetch_with(self, namespace, ttl_seconds): - if not namespace or not namespace.strip(): raise ValueError("continuous page namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("continuous page ttl_seconds must be positive") - self._continuous_page_fetch_options = {"namespace": namespace, "ttl_seconds": ttl_seconds} - return self - def optimize_pagination_with_id_set(self): - return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) - def optimize_pagination_with_id_set_config(self, namespace, ttl_seconds, max_ids): - if not namespace or not namespace.strip(): raise ValueError("ID set pagination namespace must not be empty") - if ttl_seconds <= 0: raise ValueError("ID set pagination ttl_seconds must be positive") - if max_ids <= 0: raise ValueError("ID set pagination max_ids must be positive") - self.id_set_pagination = {"namespace": namespace, "ttl_seconds": ttl_seconds, "max_ids": max_ids} - return self - -class QueryRequest: - def __init__(self, query): - self.query = query - -async def _execute_facets(service, context, outer_query): - facets = {} - for facet in getattr(outer_query, "_facets", []): - membership = copy.deepcopy(outer_query) - membership._facets = [] - membership._relations = [] - membership._order_by = [] - membership._offset = None - membership._limit = None - membership._projection = [] - membership._aggregates = [("count", "id", "__teaql_facet_count")] - membership._group_by = [facet["relation_name"]] - membership_rows = (await service.query(context, QueryRequest(membership))).rows - counts = {str(row[facet["relation_name"]]): int(row["__teaql_facet_count"]) - for row in membership_rows if row.get(facet["relation_name"]) is not None} - - nested = copy.deepcopy(facet["query"]) - nested._facets = [] - aliases = [alias for function, _field, alias in nested._aggregates - if function.lower() == "count"] or ["count"] - nested._aggregates = [] - nested._group_by = [] - nested_rows = (await service.query(context, QueryRequest(nested))).rows - decorated = [] - for row in nested_rows: - count = counts.get(str(row.get("id")), 0) - if not facet["include_all_facets"] and count == 0: continue - copy_row = dict(row) - for alias in aliases: copy_row[alias] = count - decorated.append(copy_row) - facets[facet["name"]] = SmartList(decorated) - return facets - -class MutationRequest: - def __init__(self, cmd): - self.cmd = cmd - self.comment = None - -class InsertCommand: - def __init__(self, entity, payload): - self.entity = entity - self.payload = payload - -class UpdateCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - self.values = {} - - def value(self, k, v): - self.values[k] = v - -class DeleteCommand: - def __init__(self, entity, pk, expected_version=None): - self.entity = entity - self.pk = pk - self.expected_version = expected_version - -def eq(a, b): return {"type": "eq", "field": a, "value": b} -def ne(a, b): return {"type": "ne", "field": a, "value": b} -def contain(a, b): return {"type": "contain", "field": a, "value": b} -def not_contain(a, b): return {"type": "not_contain", "field": a, "value": b} -def begin_with(a, b): return {"type": "begin_with", "field": a, "value": b} -def not_begin_with(a, b): return {"type": "not_begin_with", "field": a, "value": b} -def end_with(a, b): return {"type": "end_with", "field": a, "value": b} -def not_end_with(a, b): return {"type": "not_end_with", "field": a, "value": b} -def sound_like(a, b): return {"type": "sound_like", "field": a, "value": b} -def one_of(a, values): return {"type": "in", "field": a, "value": list(values)} -def in_list(a, values): return one_of(a, values) -def not_in_list(a, values): return {"type": "not_in", "field": a, "value": list(values)} -def gte(a, b): return {"type": "gte", "field": a, "value": b} -def lte(a, b): return {"type": "lte", "field": a, "value": b} -def gt(a, b): return {"type": "gt", "field": a, "value": b} -def lt(a, b): return {"type": "lt", "field": a, "value": b} -def column(a): return a -def value(a): return a -def between(a, lower, upper): return {"type": "between", "field": a, "value": [lower, upper]} -def is_null(a): return {"type": "is_null", "field": a} -def is_not_null(a): return {"type": "is_not_null", "field": a} -def in_subquery(left, entity, query): - return {"type": "in_subquery", "field": left, "entity": entity, "query": query} -def not_in_subquery(left, entity, query): - return {"type": "not_in_subquery", "field": left, "entity": entity, "query": query} - -def _soundex(value): - text = "".join(ch for ch in str(value or "").upper() if "A" <= ch <= "Z") - if not text: return "?000" - groups = {**dict.fromkeys("BFPV", "1"), **dict.fromkeys("CGJKQSXZ", "2"), - **dict.fromkeys("DT", "3"), "L": "4", **dict.fromkeys("MN", "5"), "R": "6"} - result, previous = text[0], groups.get(text[0], "") - for char in text[1:]: - code = groups.get(char, "") - if code and code != previous: result += code - previous = code - if len(result) == 4: break - return (result + "000")[:4] - -def _prepare_continuous_page(context, original): - query = copy.deepcopy(original) - options = getattr(query, "_continuous_page_fetch_options", None) - if options is None or context is None or not hasattr(context, "continuous_page_cursor"): - return query, None - if query._limit is None or query._limit <= 0 or len(query._order_by) != 1 or query._order_by[0][0] != "id": - context.observe_continuous_page("OFFSET_FALLBACK:UNSUPPORTED_QUERY_SHAPE") - return query, None - normalized = copy.deepcopy(query) - normalized._offset = 0 - normalized._comment = None - normalized._purpose = None - normalized._continuous_page_fetch_options = None - owner = context.get_resource("user_identifier") or "" - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:continuous-page:v1:{digest}" - execution = {"query_key": query_key, "offset": query._offset or 0, "limit": query._limit, - "direction": query._order_by[0][1].lower(), "ttl": options["ttl_seconds"], "optimized": False} - if execution["offset"] == 0: - context.observe_continuous_page("OFFSET_FALLBACK:FIRST_PAGE") - return query, execution - cursor = context.continuous_page_cursor(query_key, execution["offset"]) - if cursor is None: - context.observe_continuous_page("OFFSET_FALLBACK:CACHE_MISS") - return query, execution - query._filters.append((lt if execution["direction"] == "desc" else gt)("id", cursor["boundary"])) - query._offset = 0 - execution["optimized"] = True - execution["cursor_id"] = cursor["cursor_id"] - context.observe_continuous_page("CURSOR_SEEK", cursor["cursor_id"]) - return query, execution - -def _register_continuous_page(context, execution, rows): - if execution is None or len(rows) != execution["limit"] or not rows or "id" not in rows[-1]: return - cursor_id = f"cpg_{time.time_ns():x}" - next_offset = execution["offset"] + len(rows) - context.put_continuous_page_cursor(execution["query_key"], next_offset, { - "cursor_id": cursor_id, "boundary": rows[-1]["id"], "expires_at": time.time() + execution["ttl"] - }) - if execution["optimized"]: context.observe_continuous_page("CURSOR_SEEK", execution["cursor_id"]) - -class MutationResult(dict): - def __init__(self, values, persisted_record=None): - super().__init__(values) - self.persisted_record = persisted_record - - -class TeaQLClient: - def __init__(self, storage_path=None): - self.storage_path = storage_path - self._data = {} - self._next_ids = {} - self._graph_snapshot = None - self._load() - - async def begin(self, context): - if self._graph_snapshot is not None: - raise RuntimeError("A graph transaction is already active on this data service") - self._graph_snapshot = (copy.deepcopy(self._data), copy.deepcopy(self._next_ids)) - return self - - async def commit(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._persist() - self._graph_snapshot = None - - async def rollback(self, context): - if self._graph_snapshot is None: - raise RuntimeError("No graph transaction is active") - self._data, self._next_ids = self._graph_snapshot - self._graph_snapshot = None - self._persist() - - def _load(self): - if not self.storage_path or not os.path.exists(self.storage_path): - return - with open(self.storage_path, "r", encoding="utf-8") as stream: - state = json.load(stream) - self._data = state.get("data", {}) - self._next_ids = state.get("next_ids", {}) - - def _persist(self): - if not self.storage_path: - return - parent = os.path.dirname(os.path.abspath(self.storage_path)) - os.makedirs(parent, exist_ok=True) - fd, temporary_path = tempfile.mkstemp(prefix=".teaql-", suffix=".json", dir=parent) - try: - with os.fdopen(fd, "w", encoding="utf-8") as stream: - json.dump({"data": self._data, "next_ids": self._next_ids}, stream) - os.replace(temporary_path, self.storage_path) - finally: - if os.path.exists(temporary_path): - os.unlink(temporary_path) - - def _next_id(self, entity): - value = int(self._next_ids.get(entity, 1)) - self._next_ids[entity] = value + 1 - return value - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - table = self._data.setdefault(command.entity, {}) - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - record_id = record.get("id") or self._next_id(command.entity) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - table[str(record_id)] = record - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "values"): - record_id = command.pk - key = str(record_id) - if key not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - record = table[key] - if command.expected_version is not None and record.get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - record.update(copy.deepcopy(command.values)) - record["version"] = int(record.get("version") or 0) + 1 - if self._graph_snapshot is None: - self._persist() - result = MutationResult( - {"success": True, "id": record_id, "version": record["version"]}, - copy.deepcopy(record)) - await context.emit_mutation_audit(req, result) - return result - if hasattr(command, "pk"): - record_id = command.pk - if str(record_id) not in table: - raise KeyError(f"{command.entity}({record_id}) does not exist") - if command.expected_version is not None and table[str(record_id)].get("version") != command.expected_version: - raise RuntimeError( - f"Optimistic lock failed for {command.entity}({record_id}): " - f"expected version {command.expected_version}" - ) - current_version = int(table[str(record_id)].get("version") or 0) - table[str(record_id)]["version"] = -(current_version + 1) - if self._graph_snapshot is None: - self._persist() - persisted = copy.deepcopy(table[str(record_id)]) - result = MutationResult({ - "success": True, "id": record_id, - "version": persisted["version"], "deleted": True, - }, persisted) - await context.emit_mutation_audit(req, result) - return result - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - - async def query(self, context, req): - query, continuous = _prepare_continuous_page(context, req.query) - rows = [copy.deepcopy(row) for row in self._data.get(query.entity, {}).values()] - for expression in query._filters: - if expression.get("type") in ("in_subquery", "not_in_subquery"): - child_result = await self.query(context, QueryRequest(expression["query"])) - projected = expression["query"]._projection - projected_field = projected[0] if projected else "id" - child_values = {row.get(projected_field) for row in child_result.rows} - if expression.get("type") == "in_subquery": - rows = [row for row in rows if row.get(expression["field"]) in child_values] - else: - rows = [row for row in rows if row.get(expression["field"]) not in child_values] - elif expression.get("type") == "eq": - rows = [row for row in rows if row.get(expression["field"]) == expression["value"]] - elif expression.get("type") == "contain": - rows = [row for row in rows if expression["value"] in str(row.get(expression["field"], ""))] - elif expression.get("type") == "not_contain": - rows = [row for row in rows if expression["value"] not in str(row.get(expression["field"], ""))] - elif expression.get("type") == "begin_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "not_begin_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).startswith(str(expression["value"]))] - elif expression.get("type") == "end_with": - rows = [row for row in rows if str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "not_end_with": - rows = [row for row in rows if not str(row.get(expression["field"], "")).endswith(str(expression["value"]))] - elif expression.get("type") == "sound_like": - rows = [row for row in rows if _soundex(row.get(expression["field"])) == _soundex(expression["value"])] - elif expression.get("type") == "in": - rows = [row for row in rows if row.get(expression["field"]) in expression["value"]] - elif expression.get("type") == "not_in": - rows = [row for row in rows if row.get(expression["field"]) not in expression["value"]] - elif expression.get("type") == "ne": - rows = [row for row in rows if row.get(expression["field"]) != expression["value"]] - elif expression.get("type") == "between": - rows = [row for row in rows if expression["value"][0] <= row.get(expression["field"]) <= expression["value"][1]] - elif expression.get("type") == "is_null": - rows = [row for row in rows if row.get(expression["field"]) is None] - elif expression.get("type") == "is_not_null": - rows = [row for row in rows if row.get(expression["field"]) is not None] - elif expression.get("type") == "gte": - rows = [row for row in rows if row.get(expression["field"]) >= expression["value"]] - elif expression.get("type") == "lte": - rows = [row for row in rows if row.get(expression["field"]) <= expression["value"]] - elif expression.get("type") == "gt": - rows = [row for row in rows if row.get(expression["field"]) > expression["value"]] - elif expression.get("type") == "lt": - rows = [row for row in rows if row.get(expression["field"]) < expression["value"]] - if query._aggregates: - if query._group_by: - grouped = {} - for row in rows: - key = tuple(row.get(field) for field in query._group_by) - grouped.setdefault(key, []).append(row) - aggregate_rows = [] - for key, group_rows in grouped.items(): - values = dict(zip(query._group_by, key)) - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(group_rows) - aggregate_rows.append(values) - return type('QueryResult', (object,), {'rows': aggregate_rows, 'facets': {}}) - values = {} - for function, _field, alias in query._aggregates: - if function.lower() != "count": raise ValueError(f"Unsupported local aggregate: {function}") - values[alias] = len(rows) - return type('QueryResult', (object,), {'rows': [values], 'facets': {}}) - for field, direction in reversed(query._order_by): - rows.sort(key=lambda row: (row.get(field) is None, row.get(field)), reverse=direction.lower() == "desc") - start = query._offset or 0 - end = None if query._limit is None else start + query._limit - result_rows = rows[start:end] - _register_continuous_page(context, continuous, result_rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': result_rows, 'facets': facets}) - - async def close(self): - pass - - -class _Transaction: - def __init__(self, connection): - self.connection = connection - - async def __aenter__(self): - await self.connection.begin() - return self.connection - - async def __aexit__(self, exc_type, exc, traceback): - if exc_type is None: - await self.connection.commit() - else: - await self.connection.rollback() - - -class _NoopTransaction: - async def __aenter__(self): return self - async def __aexit__(self, exc_type, exc, traceback): return False - - -class _AsyncSqlGraphTransaction: - def __init__(self, client, connection): - self.client, self.connection = client, connection - - async def mutate(self, context, request): - return await self.client.mutate(context, request) - - async def query(self, context, request): - return await self.client.query(context, request) - - async def commit(self, context): - try: - await self.connection.commit() - finally: - await self.connection.close() - self.client._graph_connection = None - - async def rollback(self, context): - try: - await self.connection.rollback() - finally: - await self.connection.close() - self.client._graph_connection = None - - -class _PostgreSQLConnection: - def __init__(self, raw): - self.raw = raw - self.current_transaction = None - - def transaction(self): return _Transaction(self) - async def begin(self): - self.current_transaction = self.raw.transaction() - await self.current_transaction.start() - async def commit(self): - await self.current_transaction.commit() - self.current_transaction = None - async def rollback(self): - await self.current_transaction.rollback() - self.current_transaction = None - async def execute(self, sql, *params): - status = await self.raw.execute(sql, *params) - try: return int(status.rsplit(" ", 1)[-1]) - except ValueError: return -1 - async def fetch_all(self, sql, *params): - return [dict(row) for row in await self.raw.fetch(sql, *params)] - async def fetch_one(self, sql, *params): - row = await self.raw.fetchrow(sql, *params) - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - return await self.raw.fetchval(sql, *params) - async def close(self): await self.raw.close() - - -class _SQLiteConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.execute("BEGIN") - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - cursor = await self.raw.execute(sql, params) - affected = cursor.rowcount - await cursor.close() - return affected - async def fetch_all(self, sql, *params): - cursor = await self.raw.execute(sql, params) - rows = [dict(row) for row in await cursor.fetchall()] - await cursor.close() - return rows - async def fetch_one(self, sql, *params): - cursor = await self.raw.execute(sql, params) - row = await cursor.fetchone() - await cursor.close() - return None if row is None else dict(row) - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): await self.raw.close() - - -class _MySQLConnection: - def __init__(self, raw): self.raw = raw - def transaction(self): return _Transaction(self) - async def begin(self): await self.raw.begin() - async def commit(self): await self.raw.commit() - async def rollback(self): await self.raw.rollback() - async def execute(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return cursor.rowcount - async def fetch_all(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return list(await cursor.fetchall()) - async def fetch_one(self, sql, *params): - async with self.raw.cursor() as cursor: - await cursor.execute(sql, params) - return await cursor.fetchone() - async def fetch_value(self, sql, *params): - row = await self.fetch_one(sql, *params) - return None if row is None else next(iter(row.values())) - async def close(self): self.raw.close() - - -class AsyncSqlTeaQLClient: - """Shared async SQL persistence for PostgreSQL, MySQL, and SQLite.""" - - database_kind = None - identifier_quote = '"' - _identifier_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") - _type_maps = { - "postgres": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE PRECISION", - "decimal": "NUMERIC", "date": "DATE", "datetime": "TIMESTAMPTZ", - "json": "JSONB", "text": "TEXT", - }, - "mysql": { - "bool": "BOOLEAN", "integer": "BIGINT", "float": "DOUBLE", - "decimal": "DECIMAL(38, 10)", "date": "DATE", "datetime": "DATETIME(6)", - "json": "JSON", "text": "TEXT", - }, - "sqlite": { - "bool": "INTEGER", "integer": "INTEGER", "float": "REAL", - "decimal": "NUMERIC", "date": "TEXT", "datetime": "TEXT", - "json": "TEXT", "text": "TEXT", - }, - } - - def __init__(self, database_url): - if not database_url: - raise ValueError("database_url is required") - self.database_url = database_url - self._graph_connection = None - - async def begin(self, context): - if self._graph_connection is not None: - raise RuntimeError("A graph transaction is already active on this data service") - connection = await self._connect() - await connection.begin() - self._graph_connection = connection - return _AsyncSqlGraphTransaction(self, connection) - - @staticmethod - def _table_name(entity): - schema = ENTITY_SCHEMAS.get(entity) - if schema is not None: - return schema["table"] - snake = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", entity) - snake = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", snake).lower() - return f"{snake}_data" - - def _identifier(self, value): - if not self._identifier_pattern.fullmatch(value): - raise ValueError(f"Unsafe SQL identifier: {value!r}") - quote = self.identifier_quote - return f"{quote}{value}{quote}" - - def _placeholder(self, index): - if self.database_kind == "postgres": return f"${index}" - if self.database_kind == "mysql": return "%s" - return "?" - - def _normalize(self, value): - value = getattr(value, "id", value) - if isinstance(value, (dict, list)): - return json.dumps(value) - if self.database_kind == "sqlite" and isinstance(value, Decimal): - return str(value) - if self.database_kind == "sqlite" and isinstance(value, (date, datetime)): - return value.isoformat() - return value - - @staticmethod - def _logical_type(value): - value = getattr(value, "id", value) - if isinstance(value, bool): return "bool" - if isinstance(value, int): return "integer" - if isinstance(value, float): return "float" - if isinstance(value, Decimal): return "decimal" - if isinstance(value, datetime): return "datetime" - if isinstance(value, date): return "date" - if isinstance(value, (dict, list)): return "json" - return "text" - - def _column_type(self, logical_type): - return self._type_maps[self.database_kind].get(logical_type, "BIGINT") - - async def _column_exists(self, connection, table, field): - if self.database_kind == "postgres": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = current_schema() AND table_name = $1 AND column_name = $2", - table, field, - ) - return value is not None - if self.database_kind == "mysql": - value = await connection.fetch_value( - "SELECT 1 FROM information_schema.columns " - "WHERE table_schema = DATABASE() AND table_name = %s AND column_name = %s", - table, field, - ) - return value is not None - rows = await connection.fetch_all(f"PRAGMA table_info({self._identifier(table)})") - return any(row["name"] == field for row in rows) - - async def _ensure_table(self, connection, entity, values=None): - table = self._table_name(entity) - quoted_table = self._identifier(table) - await connection.execute( - f"CREATE TABLE IF NOT EXISTS {quoted_table} (" - f"{self._identifier('id')} BIGINT PRIMARY KEY, " - f"{self._identifier('version')} BIGINT NOT NULL)" - ) - columns = dict(ENTITY_SCHEMAS.get(entity, {}).get("columns", {})) - required = dict(ENTITY_SCHEMAS.get(entity, {}).get("required", {})) - for field, value in (values or {}).items(): - columns.setdefault(field, self._logical_type(value)) - for field, logical_type in columns.items(): - if field in ("id", "version") or await self._column_exists(connection, table, field): - continue - await connection.execute( - f"ALTER TABLE {quoted_table} ADD COLUMN {self._identifier(field)} " - f"{self._column_type(logical_type)}" - f"{' NOT NULL' if required.get(field, False) else ''}" - ) - return table - - async def _ensure_schema(self, context, invocation): - if invocation is not _SCHEMA_INVOCATION: - raise PermissionError("Ensure Schema must be invoked through UserContext.ensure_schema()") - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - for entity in ENTITY_SCHEMAS: - await self._ensure_table(connection, entity) - if context is not None: - roots = context.get_resource("root_graphs") or () - constants = context.get_resource("initial_graphs") or () - for graph, reconcile in (tuple((g, False) for g in roots) - + tuple((g, True) for g in constants)): - table = await self._ensure_table(connection, graph.entity, graph.fields) - seed_id = int(graph.fields["id"]) - existing = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} WHERE {self._identifier('id')} = {self._placeholder(1)}", - seed_id) - if existing is None: - record = dict(graph.fields) - record["version"] = int(record.get("version") or 1) - fields = list(record) - await connection.execute( - f"INSERT INTO {self._identifier(table)} ({', '.join(self._identifier(f) for f in fields)}) VALUES ({', '.join(self._placeholder(i) for i in range(1, len(fields)+1))})", - *(self._normalize(record[f]) for f in fields)) - elif reconcile: - existing = dict(existing) - changed = {k: v for k, v in graph.fields.items() - if k != "id" and existing.get(k) != self._normalize(v)} - if changed: - fields = list(changed) - next_index = len(fields) + 1 - await connection.execute( - f"UPDATE {self._identifier(table)} SET {', '.join(self._identifier(f) + ' = ' + self._placeholder(i) for i, f in enumerate(fields, 1))}, {self._identifier('version')} = {self._identifier('version')} + 1 WHERE {self._identifier('id')} = {self._placeholder(next_index)}", - *(self._normalize(changed[f]) for f in fields), seed_id) - await self._ensure_id_floor(connection, graph.entity, seed_id) - finally: - if owns_connection: - await connection.close() - - async def _next_id(self, connection, entity): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", 1)", entity) - return 1 - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= 2**63 - 1: - raise RuntimeError(f"ID space overflow for {entity}") - next_value = current + 1 - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - next_value, entity, current) - if changed == 1: - return next_value - if changed not in (0, None): - raise RuntimeError( - f"ID space update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to allocate ID for {entity} after 100 optimistic-lock attempts") - - async def _ensure_id_floor(self, connection, entity, floor): - await connection.execute( - "CREATE TABLE IF NOT EXISTS teaql_id_space (" - "type_name VARCHAR(255) PRIMARY KEY, current_level BIGINT NOT NULL)" - ) - for attempt in range(1, 101): - current = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if current is None: - try: - await connection.execute( - "INSERT INTO teaql_id_space(type_name, current_level) VALUES (" - + self._placeholder(1) + ", " + self._placeholder(2) + ")", - entity, floor) - return - except Exception: - winner = await connection.fetch_value( - "SELECT current_level FROM teaql_id_space WHERE type_name = " - + self._placeholder(1), entity) - if winner is None: - raise - continue - current = int(current) - if current >= floor: - return - changed = await connection.execute( - "UPDATE teaql_id_space SET current_level = " + self._placeholder(1) - + " WHERE type_name = " + self._placeholder(2) - + " AND current_level = " + self._placeholder(3), - floor, entity, current) - if changed == 1: - return - if changed not in (0, None): - raise RuntimeError( - f"ID space floor update for {entity} changed {changed} rows on attempt {attempt}") - raise RuntimeError( - f"Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts") - - async def mutate(self, context, req): - command = req.cmd - if not context.consume_mutation_checked(command): - context.check_and_fix_mutation(command) - started_ns = time.perf_counter_ns() - owns_connection = self._graph_connection is None - connection = await self._connect() if owns_connection else self._graph_connection - try: - async with (connection.transaction() if owns_connection else _NoopTransaction()): - if hasattr(command, "payload"): - record = copy.deepcopy(command.payload) - table = await self._ensure_table(connection, command.entity, record) - record_id = record.get("id") or await self._next_id(connection, command.entity) - if record.get("id") is not None: - await self._ensure_id_floor(connection, command.entity, int(record_id)) - record["id"] = record_id - record["version"] = int(record.get("version") or 0) + 1 - fields = list(record.keys()) - columns = ", ".join(self._identifier(field) for field in fields) - placeholders = ", ".join( - self._placeholder(index) for index in range(1, len(fields) + 1) - ) - params = [self._normalize(record[field]) for field in fields] - sql = f"INSERT INTO {self._identifier(table)} ({columns}) VALUES ({placeholders})" - await connection.execute(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Insert, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=1, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "insert"))) - persisted = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - record_id, - ) - result = MutationResult( - {"success": True, "id": record_id, "version": persisted["version"]}, - persisted) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "values"): - table = await self._ensure_table(connection, command.entity, command.values) - values = { - field: value for field, value in command.values.items() - if field not in ("id", "version") - } - params = [self._normalize(value) for value in values.values()] - assignments = [ - f"{self._identifier(field)} = {self._placeholder(index)}" - for index, field in enumerate(values.keys(), 1) - ] - version = self._identifier("version") - assignments.append(f"{version} = {version} + 1") - params.append(command.pk) - predicates = [ - f"{self._identifier('id')} = {self._placeholder(len(params))}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{version} = {self._placeholder(len(params))}" - ) - sql = (f"UPDATE {self._identifier(table)} SET {', '.join(assignments)} " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Update, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "update"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult( - {"success": True, "id": command.pk, "version": row["version"]}, row) - await context.emit_mutation_audit(req, result) - return result - - if hasattr(command, "pk"): - table = await self._ensure_table(connection, command.entity) - params = [command.pk] - predicates = [ - f"{self._identifier('id')} = {self._placeholder(1)}" - ] - if command.expected_version is not None: - params.append(command.expected_version) - predicates.append( - f"{self._identifier('version')} = {self._placeholder(len(params))}" - ) - version = self._identifier("version") - sql = (f"UPDATE {self._identifier(table)} SET {version} = -({version} + 1) " - f"WHERE {' AND '.join(predicates)}") - affected = await connection.execute(sql, *params) - if affected != 1: - raise RuntimeError( - f"Optimistic lock failed or {command.entity}({command.pk}) does not exist" - ) - context.record_sql_evidence( - SqlLogOperation.Delete, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, affected_rows=affected, - audit_reason=req.comment, - trace_path=(("operation", "mutation"), ("entity", command.entity), - ("provider", self.database_kind), ("sql", "delete"))) - row = await connection.fetch_one( - f"SELECT * FROM {self._identifier(table)} " - f"WHERE {self._identifier('id')} = {self._placeholder(1)}", - command.pk, - ) - result = MutationResult({ - "success": True, "id": command.pk, - "version": row["version"], "deleted": True, - }, row) - await context.emit_mutation_audit(req, result) - return result - - raise TypeError(f"Unsupported mutation command: {type(command).__name__}") - finally: - if owns_connection: - await connection.close() - - def _contains_predicate(self, field, placeholder): - if self.database_kind == "mysql": - return f"CAST({field} AS CHAR) LIKE CONCAT('%%', {placeholder}, '%%')" - return f"CAST({field} AS TEXT) LIKE '%' || {placeholder} || '%'" - - def _compile_filter_expression(self, expression, params): - field = self._identifier(expression["field"]) - operator = expression.get("type") - if operator in ("in_subquery", "not_in_subquery"): - child = expression["query"] - projection = child._projection[0] if child._projection else "id" - projected = self._identifier(projection) - child_predicates = [ - self._compile_filter_expression(item, params) for item in child._filters - ] - child_schema = ENTITY_SCHEMAS.get(child.entity, {}) - if "version" in child_schema.get("columns", {}): - child_predicates.append(f"{self._identifier('version')} > 0") - negative = operator == "not_in_subquery" - if negative: - child_predicates.append(f"{projected} IS NOT NULL") - where = " WHERE " + " AND ".join(child_predicates) if child_predicates else "" - child_sql = (f"SELECT {projected} FROM " - f"{self._identifier(self._table_name(child.entity))}{where}") - return f"{field} {'NOT IN' if negative else 'IN'} ({child_sql})" - if operator in ("in", "not_in"): - values = list(expression.get("value") or []) - if not values: - return "1 = 0" if operator == "in" else "1 = 1" - placeholders = [] - for value in values: - params.append(self._normalize(value)) - placeholders.append(self._placeholder(len(params))) - return f"{field} {'IN' if operator == 'in' else 'NOT IN'} ({', '.join(placeholders)})" - if operator in ("is_null", "is_not_null"): - return f"{field} IS {'NULL' if operator == 'is_null' else 'NOT NULL'}" - if operator == "between": - bounds = list(expression.get("value") or []) - if len(bounds) != 2: - raise ValueError("between requires exactly two bounds") - params.extend([self._normalize(bounds[0]), self._normalize(bounds[1])]) - return (f"{field} BETWEEN {self._placeholder(len(params)-1)} " - f"AND {self._placeholder(len(params))}") - if operator == "sound_like": - params.append(self._normalize(expression.get("value"))) - return f"SOUNDEX({field}) = SOUNDEX({self._placeholder(len(params))})" - raw_value = expression.get("value") - params.append(self._normalize(raw_value)) - placeholder = self._placeholder(len(params)) - if operator == "eq": return f"{field} = {placeholder}" - if operator == "ne": return f"{field} <> {placeholder}" - if operator == "contain": return self._contains_predicate(field, placeholder) - if operator == "not_contain": return f"NOT ({self._contains_predicate(field, placeholder)})" - if operator in ("begin_with", "not_begin_with", "end_with", "not_end_with"): - raw = str(raw_value or "") - params[-1] = ("%" if "end" in operator else "") + raw + ("%" if "begin" in operator else "") - clause = f"{field} LIKE {placeholder}" - return f"NOT ({clause})" if operator.startswith("not_") else clause - if operator == "gte": return f"{field} >= {placeholder}" - if operator == "lte": return f"{field} <= {placeholder}" - if operator == "gt": return f"{field} > {placeholder}" - if operator == "lt": return f"{field} < {placeholder}" - params.pop() - raise ValueError(f"Unsupported filter operator: {operator}") - - async def _prepare_id_set_page(self, context, original): - query = copy.deepcopy(original) - options = getattr(query, "id_set_pagination", None) - if options is None or context is None or not hasattr(context, "id_set_get"): - if context is not None and hasattr(context, "observe_id_set"): - context.observe_id_set("ID_SET_DISABLED") - return query, [], False - if query._limit is None or query._limit <= 0 or query._partition_by is not None or query._aggregates or query._group_by: - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - stable = copy.deepcopy(query) - if not any(field == "id" for field, _direction in stable._order_by): - stable._order_by.append(("id", "asc")) - normalized = copy.deepcopy(stable) - normalized._offset = None; normalized._limit = None - normalized._projection = []; normalized._relations = []; normalized._relation_aggregates = [] - normalized._facets = []; normalized._comment = None; normalized._purpose = None - normalized.id_set_pagination = None - owner = context.get_resource("user_identifier") or "" - active_root = context.get_resource("active_root") - policy = context.get_resource("request_policy") - source = context.get_resource("dataService") - digest = hashlib.sha256( - f'{options["namespace"]}|{owner}|{id(source)}|{id(policy)}|{active_root!r}|{vars(normalized)!r}'.encode("utf-8") - ).hexdigest() - query_key = f"teaql:id-set:v1:{digest}" - retained = context.id_set_get(query_key) - plan = "ID_SET_HIT" - if retained is None: - async with context.id_set_lock(query_key): - retained = context.id_set_get(query_key) - if retained is None: - id_query = copy.deepcopy(stable) - id_query._projection = ["id"] - id_query._relations = []; id_query._relation_aggregates = []; id_query._facets = [] - id_query._offset = 0; id_query._limit = options["max_ids"] + 1 - id_query.id_set_pagination = None - id_rows = (await self.query(context, QueryRequest(id_query))).rows - try: ids = tuple(int(row["id"]) for row in id_rows) - except (KeyError, TypeError, ValueError): - context.observe_id_set("ID_SET_FALLBACK_UNSUPPORTED_SHAPE") - return query, [], False - if len(ids) > options["max_ids"]: - context.observe_id_set("ID_SET_FALLBACK_LIMIT_EXCEEDED", "LOWER_BOUND", len(ids)) - return query, [], False - try: context.id_set_put(query_key, ids, options["ttl_seconds"]) - except Exception: - context.observe_id_set("ID_SET_FALLBACK_STORE_UNAVAILABLE") - return query, [], False - retained = context.id_set_get(query_key) - plan = "ID_SET_BUILD" - ids = retained["ids"] - context.observe_id_set(plan, "EXACT", len(ids)) - start = query._offset or 0 - if start >= len(ids): return query, [], True - page_ids = list(ids[start:min(start + query._limit, len(ids))]) - query._offset = None; query._limit = None; query.id_set_pagination = None - query._filters.append(in_list("id", page_ids)) - return query, page_ids, False - - async def query(self, context, req): - started_ns = time.perf_counter_ns() - query, id_set_order, id_set_empty = await self._prepare_id_set_page(context, req.query) - if id_set_empty: - return type('QueryResult', (object,), {'rows': [], 'facets': {}}) - query, continuous = _prepare_continuous_page(context, query) - filter_values = { - expression["field"]: expression.get("value") for expression in query._filters - } - connection = await self._connect() - try: - table = await self._ensure_table(connection, query.entity, filter_values) - params = [] - predicates = [] - for expression in query._filters: - predicates.append(self._compile_filter_expression(expression, params)) - - group_fields = [self._identifier(field) for field in query._group_by] - if query._aggregates: - projections = list(group_fields) - functions = { - "count": "COUNT", "sum": "SUM", "avg": "AVG", - "min": "MIN", "max": "MAX", "stddev": "STDDEV", - "stddev_pop": "STDDEV_POP", "var_samp": "VAR_SAMP", - "var_pop": "VAR_POP", "bit_and": "BIT_AND", - "bit_or": "BIT_OR", "bit_xor": "BIT_XOR", - } - for function, field, alias in query._aggregates: - sql_function = functions.get(function.lower()) - if sql_function is None: - raise ValueError(f"Unsupported aggregate function: {function}") - projections.append( - f"{sql_function}({self._identifier(field)}) AS {self._identifier(alias)}" - ) - projection = ", ".join(projections) - else: - projection = ", ".join(self._identifier(field) for field in query._projection) if query._projection else "*" - - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - partition_by = getattr(query, "_partition_by", None) - if partition_by: - window_order = "" - if query._order_by: - window_orders = [] - for order_field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - window_orders.append(f"{self._identifier(order_field)} {normalized_direction}") - window_order = " ORDER BY " + ", ".join(window_orders) - projection += ( - f", ROW_NUMBER() OVER (PARTITION BY {self._identifier(partition_by)}" - f"{window_order}) AS {self._identifier('__teaql_partition_rank')}" - ) - sql = f"SELECT {projection} FROM {self._identifier(table)}" - if predicates: sql += " WHERE " + " AND ".join(predicates) - if group_fields: sql += " GROUP BY " + ", ".join(group_fields) - - if query._order_by and not partition_by: - orders = [] - for field, direction in query._order_by: - normalized_direction = direction.upper() - if normalized_direction not in ("ASC", "DESC"): - raise ValueError(f"Unsupported order direction: {direction}") - orders.append(f"{self._identifier(field)} {normalized_direction}") - sql += " ORDER BY " + ", ".join(orders) - if partition_by: - rank = self._identifier("__teaql_partition_rank") - rank_predicates = [] - params.append(int(query._offset or 0)) - rank_predicates.append(f"{rank} > {self._placeholder(len(params))}") - if query._limit is not None: - params.append(int(query._offset or 0) + int(query._limit)) - rank_predicates.append(f"{rank} <= {self._placeholder(len(params))}") - sql = (f"SELECT * FROM ({sql}) AS {self._identifier('__teaql_partitioned')} " - f"WHERE {' AND '.join(rank_predicates)} ORDER BY {rank}") - elif query._limit is not None: - params.append(int(query._limit)) - sql += f" LIMIT {self._placeholder(len(params))}" - elif query._offset is not None and self.database_kind == "sqlite": - sql += " LIMIT -1" - elif query._offset is not None and self.database_kind == "mysql": - sql += " LIMIT 18446744073709551615" - if query._offset is not None and not partition_by: - params.append(int(query._offset)) - sql += f" OFFSET {self._placeholder(len(params))}" - rows = await connection.fetch_all(sql, *params) - context.record_sql_evidence( - SqlLogOperation.Select, sql, params, - (time.perf_counter_ns() - started_ns) // 1000, result_count=len(rows), - comment=query._comment, purpose=query._purpose, - trace_path=(("operation", "query"), ("request", query.entity), - *query._trace_path, - ("provider", self.database_kind), ("sql", "select"))) - finally: - await connection.close() - - await self._enhance_relations(context, query, rows) - await self._enhance_relation_aggregates(context, query, rows) - if id_set_order: - by_id = {int(row["id"]): row for row in rows if row.get("id") is not None} - rows = [by_id[entity_id] for entity_id in id_set_order if entity_id in by_id] - _register_continuous_page(context, continuous, rows) - facets = await _execute_facets(self, context, query) - return type('QueryResult', (object,), {'rows': rows, 'facets': facets}) - - async def _enhance_relations(self, context, query, parents): - if not parents or not getattr(query, "_relations", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for load in query._relations: - relation = relations.get(load["name"]) - if relation is None: raise ValueError(f"Missing relation {query.entity}.{load['name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child_query = copy.deepcopy(load["query"]) - child_query._comment = query._comment - child_query._purpose = query._purpose - child_query._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{load['name']}")] - child_query._continuous_page_fetch_options = None - child_query.entity = relation["target_entity"] - if relation["foreign_key"] not in child_query._projection: - child_query._projection.append(relation["foreign_key"]) - child_query._filters.append(one_of(relation["foreign_key"], parent_ids)) - if child_query._limit is not None: child_query._partition_by = relation["foreign_key"] - children = (await self.query(context, QueryRequest(child_query))).rows - buckets = {} - for child in children: - child.pop("__teaql_partition_rank", None) - buckets.setdefault(child.get(relation["foreign_key"]), []).append(child) - for parent in parents: - related = buckets.get(parent.get(relation["local_key"]), []) - parent[load["name"]] = related if relation["many"] else (related[0] if related else None) - - async def _enhance_relation_aggregates(self, context, query, parents): - if not parents or not getattr(query, "_relation_aggregates", None): return - relations = ENTITY_SCHEMAS.get(query.entity, {}).get("relations", {}) - for aggregate in query._relation_aggregates: - relation = relations.get(aggregate["relation_name"]) - if relation is None: - raise ValueError(f"Missing relation {query.entity}.{aggregate['relation_name']}") - parent_ids = [p[relation["local_key"]] for p in parents if relation["local_key"] in p] - child = copy.deepcopy(aggregate["query"]) - child._comment = query._comment - child._purpose = query._purpose - child._trace_path = [*query._trace_path, - ("relation", f"{query.entity}.{aggregate['relation_name']}")] - child._continuous_page_fetch_options = None - child.entity = relation["target_entity"] - child._projection = []; child._order_by = []; child._limit = None; child._offset = None - child._relations = []; child._relation_aggregates = [] - if not child._aggregates: child._aggregates = [("count", "id", aggregate["alias"])] - if relation["foreign_key"] not in child._group_by: child._group_by.append(relation["foreign_key"]) - child._filters.append(one_of(relation["foreign_key"], parent_ids)) - rows = (await self.query(context, QueryRequest(child))).rows - buckets = {row[relation["foreign_key"]]: row for row in rows if relation["foreign_key"] in row} - is_count = (not aggregate["query"]._aggregates or - aggregate["query"]._aggregates[0][0].lower() == "count") - for parent in parents: - row = buckets.get(parent.get(relation["local_key"])) - if row is None: - parent[aggregate["alias"]] = (0 if aggregate["single_result"] and is_count - else None if aggregate["single_result"] else {}) - elif aggregate["single_result"]: - parent[aggregate["alias"]] = row.get(child._aggregates[0][2]) - else: - parent[aggregate["alias"]] = { - key: value for key, value in row.items() - if key != relation["foreign_key"]} - - async def close(self): pass - - -class PostgreSQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "postgres" - - async def _connect(self): - try: import asyncpg - except ImportError as error: - raise RuntimeError("PostgreSQL support requires asyncpg") from error - return _PostgreSQLConnection(await asyncpg.connect(self.database_url)) - - -class MySQLTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "mysql" - identifier_quote = "`" - - async def _connect(self): - try: import aiomysql - except ImportError as error: - raise RuntimeError("MySQL support requires aiomysql") from error - parsed = urlparse(self.database_url) - if parsed.scheme not in ("mysql", "mysql+aiomysql"): - raise ValueError("MySQL database_url must use mysql://") - options = parse_qs(parsed.query) - raw = await aiomysql.connect( - host=parsed.hostname or "localhost", - port=parsed.port or 3306, - user=unquote(parsed.username or ""), - password=unquote(parsed.password or ""), - db=parsed.path.lstrip("/"), - charset=options.get("charset", ["utf8mb4"])[0], - autocommit=True, - cursorclass=aiomysql.DictCursor, - ) - return _MySQLConnection(raw) - - -class SQLiteTeaQLClient(AsyncSqlTeaQLClient): - database_kind = "sqlite" - - def __init__(self, database_url): - super().__init__(database_url) - self._soundex_enabled = False - - async def _ensure_schema(self, context, invocation): - self._soundex_enabled = True - return await super()._ensure_schema(context, invocation) - - async def _connect(self): - try: import aiosqlite - except ImportError as error: - raise RuntimeError("SQLite support requires aiosqlite") from error - database = self.database_url - if database.startswith("sqlite:"): - parsed = urlparse(database) - database = parsed.path - if database == "/:memory:": database = ":memory:" - raw = await aiosqlite.connect(database, isolation_level=None) - raw.row_factory = aiosqlite.Row - if self._soundex_enabled: - await raw.create_function("soundex", 1, _soundex, deterministic=True) - await raw.execute("PRAGMA foreign_keys = ON") - return _SQLiteConnection(raw) \ No newline at end of file diff --git a/examples/school-management/teaql/runtime.py b/examples/school-management/teaql/runtime.py deleted file mode 100644 index 7f18e86..0000000 --- a/examples/school-management/teaql/runtime.py +++ /dev/null @@ -1,645 +0,0 @@ -from dataclasses import dataclass -from datetime import timedelta -from enum import Enum -import asyncio -import builtins -import contextvars -import time - -_ID_SET_STORE = {} -_ID_SET_LOCKS = {} - -_SCHEMA_INVOCATION = object() - -_CHECK_MESSAGES = { - "en": { - "required": "{location} is required", - "min": "{location} is below the minimum", - "max": "{location} exceeds the maximum", - "min_length": "{location} is too short", - "max_length": "{location} is too long", - }, - "zh-CN": { - "required": "{location} 为必填项", - "min": "{location} 小于最小值", - "max": "{location} 超过最大值", - "min_length": "{location} 长度不足", - "max_length": "{location} 长度过长", - }, -} -_SUPPORTED_LOCALES = {"en", "zh-CN", "zh-TW", "ja", "ko", "de", "fr", "es", "pt", "ar", "th", "id", "fil", "uk", "vi"} -_LOCALE_ALIASES = {"zh": "zh-CN", "zh-hans": "zh-CN", "cn": "zh-CN", "zh-hant": "zh-TW", "tw": "zh-TW", "en-us": "en", "en-gb": "en"} - -@dataclass(frozen=True) -class ObjectLocation: - segments: tuple = () - - @classmethod - def root(cls): - return cls() - - def property(self, name): - if not isinstance(name, str) or not name: - raise ValueError("A canonical KSML property name is required") - return ObjectLocation(self.segments + (("property", name),)) - - def index(self, value): - if value < 0: - raise ValueError("Object location index must not be negative") - return ObjectLocation(self.segments + (("index", value),)) - - def prefixed_by(self, prefix): - return ObjectLocation(prefix.segments + self.segments) - - @builtins.property - def model_path(self): - result = "" - for kind, value in self.segments: - result += f"[{value}]" if kind == "index" else ("." if result else "") + value - return result - - @builtins.property - def native_path(self): - # Python's generated API uses canonical snake_case property names. - return self.model_path - - @builtins.property - def instance_path(self): - def lower_camel(value): - parts = value.split("_") - return parts[0] + "".join(part[:1].upper() + part[1:] for part in parts[1:]) - def escape(value): - return str(value).replace("~", "~0").replace("/", "~1") - return "".join("/" + escape(lower_camel(value) if kind == "property" else value) - for kind, value in self.segments) - - def __str__(self): - return self.native_path - -@dataclass -class CheckResult: - rule_id: str - location: object - input_value: object = None - system_value: object = None - message: str = None - -class CheckException(Exception): - def __init__(self, violations): - self.violations = list(violations) - super().__init__("Check failed: " + "; ".join( - result.message or f"{result.rule_id}:{result.location}" - for result in self.violations)) - -@dataclass(frozen=True) -class ContextEntityRef: - entity: str - id: int - -@dataclass(frozen=True) -class FixEvidence: - entity_type: str - model_path: str - source: str - source_label: str - -class ContextRootError(Exception): - def __init__(self, reason, expected_type, active_root=None): - self.reason, self.expected_type, self.active_root = reason, expected_type, active_root - super().__init__(f"context root {reason}: expected {expected_type}") - -@dataclass(frozen=True, order=True) -class EntityKey: - entity: str - id: object - -class EntityChangeSet: - def __init__(self): self._changes = {} - def set(self, key, field, value): self._changes.setdefault(key, {})[field] = value - def changes(self): return tuple((key, dict(values)) for key, values in self._changes.items()) - def clear_entity(self, key): self._changes.pop(key, None) - def merge_from(self, other): - for key, values in other.changes(): - for field, value in values.items(): self.set(key, field, value) - def rekey(self, old_key, new_key): - values = self._changes.pop(old_key, None) - if values: self._changes.setdefault(new_key, {}).update(values) - -class EntityRoot: - def __init__(self): - self._changes = EntityChangeSet(); self._versions = {}; self._new = set(); self._deleted = set() - def current_change_set(self): return self._changes - def set(self, key, field, value): self._changes.set(key, field, value) - def mark_as_new(self, key): self._new.add(key) - def mark_as_deleted(self, key): self._changes.clear_entity(key); self._deleted.add(key) - def set_original_version(self, key, version): self._versions[key] = version - def original_version(self, key): return self._versions.get(key) - def merge_from(self, other): - if other is self: return - self._changes.merge_from(other._changes); self._versions.update(other._versions) - self._new.update(other._new); self._deleted.update(other._deleted) - def rekey(self, old_key, new_key): - self._changes.rekey(old_key, new_key) - if old_key in self._versions: self._versions[new_key] = self._versions.pop(old_key) - if old_key in self._new: self._new.remove(old_key); self._new.add(new_key) - if old_key in self._deleted: self._deleted.remove(old_key); self._deleted.add(new_key) - def clear_entity(self, key): - self._changes.clear_entity(key); self._new.discard(key); self._deleted.discard(key) - -class SqlLogOperation(str, Enum): - Select = "select" - Insert = "insert" - Update = "update" - Delete = "delete" - -@dataclass(frozen=True) -class SqlLogEntry: - operation: SqlLogOperation - comment: object - purpose: object - audit_reason: object - trace_path: tuple - sql: str - params: tuple - debug_sql: str - elapsed: timedelta - result_count: object = None - affected_rows: object = None - result_summary: str = "" - -class DiagnosticSqlLogSink: - """Value-bearing diagnostic SQL destination; the text sink is installed by default.""" - def write(self, entry): - raise NotImplementedError - -class TextDiagnosticSqlLogSink(DiagnosticSqlLogSink): - def __init__(self, writer=print): self._writer = writer - def write(self, entry): - trace = " -> ".join(f"{key}:{value}" for key, value in entry.trace_path) - comment = "" if entry.comment is None else str(entry.comment) - purpose = "" if entry.purpose is None else str(entry.purpose) - audit_reason = "" if entry.audit_reason is None else str(entry.audit_reason) - self._writer( - f"[TeaQL SQL][{entry.operation.value}][{int(entry.elapsed.total_seconds() * 1000000)}us] " - f"{entry.result_summary} comment={comment} purpose={purpose} " - f"auditReason={audit_reason} tracePath=[{trace}]\n" - f"Parameterized SQL: {entry.sql} params={entry.params!r}\n" - f"Debug SQL: {entry.debug_sql}") - -def _diagnostic_sql_literal(value): - if value is None: return "NULL" - if isinstance(value, bool): return "1" if value else "0" - if isinstance(value, (int, float)): return str(value) - if isinstance(value, (bytes, bytearray)): return "X'" + bytes(value).hex().upper() + "'" - return "'" + str(value).replace("'", "''") + "'" - -def _render_diagnostic_sql(sql, params): - rendered = sql - for value in params: - rendered = rendered.replace("?", _diagnostic_sql_literal(value), 1) - return rendered - -@dataclass(frozen=True) -class RawAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - changes: tuple - actor: str = "" - category: str = "" - -@dataclass(frozen=True) -class SafeAuditEvent: - kind: str - entity: str - entity_id: object - reason: str - fields: tuple - actor: str = "" - category: str = "" - -class UserContext: - """Runtime dependencies and trusted request state initialized by the server.""" - - def __init__(self): - self._resources = {} - self._user_identifier = "" - self._entity_root = EntityRoot() - self._standard_audit_sink = None - self._app_audit_sink = None - self._audit_policies = {} - self._entity_initializers = {} - self._managed_entities = [] - self._continuous_page_cursors = {} - self._continuous_page_plan = "DISABLED" - self._continuous_page_cursor_id = None - self._id_set_plan = "ID_SET_DISABLED" - self._id_set_count = 0 - self._id_set_count_accuracy = "UNKNOWN" - self._query_sql_log_enabled = True - self._mutation_sql_log_enabled = True - self._sql_logs = [] - self._resources["diagnostic_sql_log_sink"] = TextDiagnosticSqlLogSink() - self._checker_registry = {} - self._checked_mutations = set() - self._graph_save_active = False - self._graph_save_lock = asyncio.Lock() - self._graph_save_owner = contextvars.ContextVar( - f"teaql_graph_save_owner_{id(self)}", default=None) - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - - def begin_fix_evidence(self): - self._resources["fix_evidence_current"] = [] - return self - - def record_fix_evidence(self, entity_type, model_path, source, source_label): - normalized = str(source_label).lower() - if not entity_type or not model_path or source not in ("clock", "context") or not source_label or "authorization" in normalized or "cookie" in normalized or "token=" in normalized: - raise ValueError("Fix evidence must contain only safe framework provenance labels") - self._resources.setdefault("fix_evidence_current", []).append(FixEvidence(entity_type, model_path, source, source_label)) - return self - - def finish_fix_evidence(self): - self._resources["fix_evidence_last"] = tuple(self._resources.get("fix_evidence_current", ())) - self._resources.pop("fix_evidence_current", None) - return self - - def last_fix_evidence(self): - return self._resources.get("fix_evidence_last", ()) - - @classmethod - def new(cls): - return cls() - - def entity_root(self): - return self._entity_root - - def insert_resource(self, resource_type, resource): - self._resources[resource_type] = resource - return self - - def set_user_identifier(self, identifier): - self._user_identifier = "" if identifier is None else str(identifier) - - def user_identifier(self): - return self._user_identifier - - def set_locale_code(self, code): - if not isinstance(code, str) or not code.strip(): - raise ValueError(f"Unsupported locale: {code}") - normalized = code.strip().replace("_", "-") - canonical = next((value for value in _SUPPORTED_LOCALES if value.lower() == normalized.lower()), None) - canonical = canonical or _LOCALE_ALIASES.get(normalized.lower()) - if canonical is None: - raise ValueError(f"Unsupported locale: {code}") - self.insert_resource("locale", canonical) - return self - - def set_language_code(self, code): - return self.set_locale_code(code) - - def _translate_check_results(self, results): - locale = self.get_resource("locale") or "en" - messages = _CHECK_MESSAGES.get(locale, _CHECK_MESSAGES["en"]) - for result in results: - key = str(result.rule_id).lower() - if key == "min_str_len": key = "min_length" - if key == "max_str_len": key = "max_length" - template = messages.get(key) or _CHECK_MESSAGES["en"].get(key) or f"checker.{key}" - location = getattr(result.location, "native_path", str(result.location)) - result.message = template.replace("{location}", location) - return results - - async def execute_graph_save(self, work): - if self._graph_save_owner.get() is not None: - return await work() - async with self._graph_save_lock: - provider = self.require_resource("dataService") - begin = getattr(provider, "begin", None) - if not callable(begin): - raise RuntimeError("Configured dataService does not support graph transactions") - transaction = await begin(self) - owner_token = self._graph_save_owner.set(object()) - self._graph_save_active = True - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - from datetime import datetime - self.insert_resource("fix_time", datetime.now()) - self.begin_fix_evidence() - self.insert_resource("dataService", transaction) - try: - result = await work() - except BaseException: - try: - await transaction.rollback(self) - finally: - for action in reversed(self._graph_rollback_actions): - action() - raise - else: - try: - await transaction.commit(self) - except BaseException: - try: - await transaction.rollback(self) - finally: - for action in reversed(self._graph_rollback_actions): - action() - raise - for action in self._graph_commit_actions: - action() - return result - finally: - self.insert_resource("dataService", provider) - self._graph_save_active = False - self._graph_commit_actions = [] - self._graph_rollback_actions = [] - self._resources.pop("fix_time", None) - self.finish_fix_evidence() - self._graph_save_owner.reset(owner_token) - - def after_graph_commit(self, work): - if not self._graph_save_active: - raise RuntimeError("No graph save is active") - self._graph_commit_actions.append(work) - - def after_graph_rollback(self, work): - if not self._graph_save_active: - raise RuntimeError("No graph save is active") - self._graph_rollback_actions.append(work) - - def install(self, module): - """Install a passive metadata manifest; this never changes a database schema.""" - module.apply_to(self) - return self - - async def ensure_schema(self): - """Explicitly reconcile schema and generated bootstrap data.""" - provider = self.require_resource("dataService") - await provider._ensure_schema(self, _SCHEMA_INVOCATION) - for bootstrap in self.get_resource("_teaql_generated_bootstraps") or (): - await bootstrap(self) - - def check_and_fix_mutation(self, mutation): - checker = self._checker_registry.get(getattr(mutation, "entity", None)) - if checker is None: - return - raw_record = getattr(mutation, "payload", getattr(mutation, "values", None)) - if raw_record is None: - return - from datetime import datetime - from teaql.core.value import Value - record = {name: Value.from_any(value) for name, value in raw_record.items()} - owns_fix_time = self._resources.get("fix_time") is None - if owns_fix_time: - self.insert_resource("fix_time", datetime.now()) - self.begin_fix_evidence() - self.insert_resource("fix_operation", "insert" if hasattr(mutation, "payload") else "update") - results = [] - try: - checker.check_and_fix(self, record, None, results) - finally: - if owns_fix_time: - self._resources.pop("fix_time", None) - self.finish_fix_evidence() - self._resources.pop("fix_operation", None) - if results: - self._translate_check_results(results) - raise CheckException(results) - raw_record.clear() - raw_record.update({name: getattr(value, "val", value) for name, value in record.items()}) - - def mark_mutation_checked(self, mutation): - self._checked_mutations.add(id(mutation)) - - def consume_mutation_checked(self, mutation): - key = id(mutation) - if key not in self._checked_mutations: - return False - self._checked_mutations.remove(key) - return True - - def get_resource(self, resource_type): - return self._resources.get(resource_type) - - def require_resource(self, resource_type): - resource = self.get_resource(resource_type) - if resource is None: - raise RuntimeError(f"Required UserContext resource is missing: {resource_type}") - return resource - - def with_active_root(self, root): - if not isinstance(root, ContextEntityRef): - raise TypeError("active root must be ContextEntityRef") - return self.insert_resource("active_root", root) - - def require_active_root(self, expected_type): - root = self.get_resource("active_root") - if not isinstance(root, ContextEntityRef): - raise ContextRootError("missing", expected_type) - if root.entity != expected_type: - raise ContextRootError("type_mismatch", expected_type, root) - return root - - def with_request_policy(self, policy): - self.insert_resource("request_policy", policy) - return self - - def prepare_query(self, query): - policy = self.get_resource("request_policy") - if policy is None: return query - if callable(policy): prepared = policy(query) - elif hasattr(policy, "apply"): prepared = policy.apply(query) - else: raise TypeError("request_policy must be callable or expose apply(query)") - return query if prepared is None else prepared - - def register_entity_initializer(self, entity_name, initializer): - if not isinstance(entity_name, str) or not entity_name.strip() or not callable(initializer): - raise ValueError("entity_name and callable initializer are required") - self._entity_initializers.setdefault(entity_name, []).append(initializer) - return self - - def initialize_entity(self, entity_name, entity): - if not isinstance(entity_name, str) or not entity_name.strip() or entity is None: - raise ValueError("entity_name and entity are required") - for initializer in self._entity_initializers.get("*", ()): - initializer(self, entity) - for initializer in self._entity_initializers.get(entity_name, ()): - initializer(self, entity) - self._managed_entities.append(entity) - return entity - - def managed_entities(self): - return list(self._managed_entities) - - def continuous_page_cursor(self, query_key, offset): - cursor = self._continuous_page_cursors.get((query_key, offset)) - if cursor is not None and cursor["expires_at"] <= __import__("time").time(): - self._continuous_page_cursors.pop((query_key, offset), None) - return None - return cursor - - def put_continuous_page_cursor(self, query_key, offset, cursor): - if len(self._continuous_page_cursors) >= 4096: - oldest = min(self._continuous_page_cursors, - key=lambda key: self._continuous_page_cursors[key]["expires_at"]) - self._continuous_page_cursors.pop(oldest, None) - self._continuous_page_cursors[(query_key, offset)] = cursor - - def observe_continuous_page(self, plan, cursor_id=None): - self._continuous_page_plan = plan - self._continuous_page_cursor_id = cursor_id - - def continuous_page_plan(self): return self._continuous_page_plan - def continuous_page_cursor_id(self): return self._continuous_page_cursor_id - - def id_set_get(self, key): - retained = _ID_SET_STORE.get(key) - if retained is not None and retained["expires_at"] <= time.time(): - _ID_SET_STORE.pop(key, None) - return None - return retained - - def id_set_put(self, key, ids, ttl_seconds): - if len(ids) * 8 > 256 * 1024 * 1024: - raise ValueError("retained ID set exceeds store memory ceiling") - while len(_ID_SET_STORE) >= 64: - oldest = min(_ID_SET_STORE, key=lambda item: _ID_SET_STORE[item]["expires_at"]) - _ID_SET_STORE.pop(oldest, None) - _ID_SET_STORE[key] = {"ids": tuple(ids), "expires_at": time.time() + ttl_seconds} - - def id_set_lock(self, key): - return _ID_SET_LOCKS.setdefault((id(asyncio.get_running_loop()), key), asyncio.Lock()) - - def observe_id_set(self, plan, accuracy="UNKNOWN", count=0): - self._id_set_plan, self._id_set_count_accuracy, self._id_set_count = plan, accuracy, count - - def id_set_plan(self): return self._id_set_plan - def id_set_count(self): return self._id_set_count, self._id_set_count_accuracy - - def _set_sql_log_mode(self, mode): - self._query_sql_log_enabled = mode in ("all", "select") - self._mutation_sql_log_enabled = mode in ("all", "mutation") - self._sql_logs = [] - return self - - def enable_all_sql_log(self): return self._set_sql_log_mode("all") - def enable_select_sql_log(self): return self._set_sql_log_mode("select") - def enable_mutation_sql_log(self): return self._set_sql_log_mode("mutation") - def disable_sql_log(self): return self._set_sql_log_mode("disabled") - def disable_select_sql_log(self): - self._query_sql_log_enabled = False - return self - def disable_mutation_sql_log(self): - self._mutation_sql_log_enabled = False - return self - def clear_sql_logs(self): self._sql_logs = [] - def sql_logs(self): return list(self._sql_logs) - def with_diagnostic_sql_log_sink(self, sink): - self._resources["diagnostic_sql_log_sink"] = sink - return self - def set_diagnostic_sql_log_sink(self, sink): - self._resources["diagnostic_sql_log_sink"] = sink - - def record_sql_evidence(self, operation, sql, params, elapsed_micros, - result_count=None, affected_rows=None, comment=None, - purpose=None, audit_reason=None, trace_path=()): - is_select = operation == SqlLogOperation.Select - if ((is_select and not self._query_sql_log_enabled) - or (not is_select and not self._mutation_sql_log_enabled)): - return - summary = (f"{result_count} rows returned" if result_count is not None - else f"{affected_rows} rows affected") - entry = SqlLogEntry(operation, comment, purpose, audit_reason, tuple(trace_path), - sql, tuple(params), _render_diagnostic_sql(sql, params), - timedelta(microseconds=elapsed_micros), result_count, affected_rows, summary) - self._sql_logs.append(entry) - sink = self._resources.get("diagnostic_sql_log_sink") - if sink is not None: sink.write(entry) - - def initialize_audit(self, standard_sink, app_sink=None): - self._standard_audit_sink = standard_sink - self._app_audit_sink = app_sink - return self - - def configure_audit_policy(self, entity, mask_fields=(), max_length=None): - self._audit_policies[entity] = (frozenset(mask_fields), max_length) - return self - - async def emit_mutation_audit(self, req, result): - command = req.cmd - values = getattr(command, "payload", getattr(command, "values", {})) - kind = "created" if hasattr(command, "payload") else "updated" if hasattr(command, "values") else "deleted" - raw = RawAuditEvent( - kind, command.entity, result.get("id"), req.comment, - tuple((name, None, value) for name, value in values.items()), - self.user_identifier(), self.get_resource("bootstrapCategory") or "", - ) - if self._standard_audit_sink is not None: - emitted = self._standard_audit_sink.on_event(self, raw) - if hasattr(emitted, "__await__"): await emitted - if self._app_audit_sink is not None: - masks, limit = self._audit_policies.get(command.entity, (frozenset(), None)) - fields = [] - for name, _, raw_value in raw.changes: - value = None if raw_value is None else str(raw_value) - masked = name in masks - if value is not None and masked: - value = "*" * len(value) if len(value) < 8 else value[:2] + "*" * (len(value) - 4) + value[-2:] - truncated = value is not None and limit is not None and len(value) > limit - if truncated: value = "*" * limit if limit <= 3 else value[:limit - 3] + "..." - fields.append((name, value, masked, truncated)) - safe = SafeAuditEvent( - kind, command.entity, result.get("id"), req.comment, tuple(fields), - raw.actor, raw.category, - ) - emitted = self._app_audit_sink.on_safe_event(self, safe) - if hasattr(emitted, "__await__"): await emitted - -class RuntimeModule: - """Immutable generated runtime manifest.""" - def __init__(self, entities=(), schemas=None, checkers=None, root_graphs=(), initial_graphs=(), generated_bootstraps=()): - self.entities = tuple(entities) - self.schemas = dict(schemas or {}) - self.checkers = dict(checkers or {}) - self.root_graphs = tuple(root_graphs) - self.initial_graphs = tuple(initial_graphs) - self.generated_bootstraps = tuple(generated_bootstraps) - - def entity(self, entity): - self.entities = (*self.entities, entity) - return self - - def checker(self, entity, checker): - self.checkers[entity] = checker - return self - - def root_graph(self, graph): - self.root_graphs = (*self.root_graphs, graph) - return self - - def initial_graph(self, graph): - self.initial_graphs = (*self.initial_graphs, graph) - return self - - def generated_bootstrap(self, bootstrap): - self.generated_bootstraps = (*self.generated_bootstraps, bootstrap) - return self - - def and_module(self, other): - return RuntimeModule(self.entities + other.entities, - {**self.schemas, **other.schemas}, - {**self.checkers, **other.checkers}, - self.root_graphs + other.root_graphs, - self.initial_graphs + other.initial_graphs, - self.generated_bootstraps + other.generated_bootstraps) - - def apply_to(self, context): - context.insert_resource("entities", self.entities) - context.insert_resource("entity_schemas", dict(self.schemas)) - context._checker_registry.update(self.checkers) - context.insert_resource("root_graphs", self.root_graphs) - context.insert_resource("initial_graphs", self.initial_graphs) - context.insert_resource("_teaql_generated_bootstraps", self.generated_bootstraps) \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 882a534..13ac72e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "teaql" -version = "0.2.4" +version = "0.2.5" description = "TeaQL Runtime and Data Service for Python" readme = "README.md" requires-python = ">=3.10" diff --git a/scripts/verify-examples.sh b/scripts/verify-examples.sh index 888dea1..c87a3b4 100755 --- a/scripts/verify-examples.sh +++ b/scripts/verify-examples.sh @@ -9,10 +9,24 @@ if [[ "${actual[*]}" != "${expected[*]}" ]]; then exit 1 fi +mapfile -t embedded_runtime_dirs < <(find "$repo/examples" -type d -name teaql -print) +if ((${#embedded_runtime_dirs[@]})); then + printf 'generated examples must depend on the packaged runtime; embedded teaql snapshot: %s\n' \ + "${embedded_runtime_dirs[@]}" >&2 + exit 1 +fi +if rg -l 'include\s*=.*teaql\*' "$repo/examples" --glob pyproject.toml >/dev/null; then + echo 'generated example package discovery must not include an embedded teaql namespace' >&2 + exit 1 +fi + PYTHONPATH="$repo/examples/conformance:$repo/src" python -m app.main PYTHONPATH="$repo/examples/school-management:$repo/src" python -m app.main -PYTHONPATH="$repo/examples/order-management/python-lib-core:$repo/src" python "$repo/examples/order-management/python-app-console/app.py" +order_management_tmp="$(mktemp -d)" task_board_tmp="$(mktemp -d)" -trap 'rm -rf "$task_board_tmp"' EXIT +trap 'rm -rf "$order_management_tmp" "$task_board_tmp"' EXIT +TEAQL_ORDER_MANAGEMENT_DB="$order_management_tmp/order.db" \ + PYTHONPATH="$repo/examples/order-management/python-lib-core:$repo/src" \ + python "$repo/examples/order-management/python-app-console/app.py" TEAQL_TASK_BOARD_DB="$task_board_tmp/task_board.db" PYTHONPATH="$repo/examples/task_board:$repo/src" python "$repo/examples/task_board/main.py" echo "PASS: all Python examples" diff --git a/src/teaql/core/meta.py b/src/teaql/core/meta.py index f4960af..630d570 100644 --- a/src/teaql/core/meta.py +++ b/src/teaql/core/meta.py @@ -5,6 +5,7 @@ def __init__(self, name: str, property_type: str = "String"): self.column_name_val = name self._is_id = False self._is_version = False + self.nullable = True def column_name(self, name): self.column_name_val = name return self @@ -14,6 +15,9 @@ def is_id(self): def is_version(self): self._is_version = True return self + def required(self): + self.nullable = False + return self class RelationDescriptor: def __init__(self, name: str, target_entity: str): diff --git a/src/teaql/core/query.py b/src/teaql/core/query.py index 1c077eb..f63c6e4 100644 --- a/src/teaql/core/query.py +++ b/src/teaql/core/query.py @@ -226,6 +226,16 @@ def top_n_probe_parent_threshold(self, parent_count: int) -> 'SelectQuery': def optimize_pagination_with_id_set(self) -> 'SelectQuery': return self.optimize_pagination_with_id_set_config("default", 600, 3_000_000) + def optimize_for_continuous_page_fetch(self) -> 'SelectQuery': + return self.optimize_for_continuous_page_fetch_with("default", 600) + + def optimize_for_continuous_page_fetch_with( + self, namespace: str, ttl_seconds: int) -> 'SelectQuery': + if not namespace or ttl_seconds <= 0: + raise ValueError("continuous page namespace and positive ttl are required") + self.continuous_page_fetch = (namespace, ttl_seconds) + return self + def optimize_pagination_with_id_set_config( self, namespace: str, ttl_seconds: int, max_ids: int) -> 'SelectQuery': if not isinstance(namespace, str) or not namespace.strip(): @@ -260,6 +270,12 @@ def comment(self, text: str) -> 'SelectQuery': self.comment_text = text return self + def purpose(self, text: str) -> 'SelectQuery': + # Purpose is carried by generated request wrappers; retaining it here + # keeps the query object self-describing for diagnostics. + self.purpose_text = text + return self + def project(self, *fields: str) -> 'SelectQuery': self.projection.extend(fields) return self @@ -305,7 +321,14 @@ def or_having(self, expr: Expr) -> 'SelectQuery': self.having_expr = expr return self - def order_by(self, order: OrderBy) -> 'SelectQuery': + def order_by(self, order: OrderBy | str, direction: str | None = None) -> 'SelectQuery': + if isinstance(order, str): + normalized = (direction or "asc").lower() + if normalized not in ("asc", "desc"): + raise ValueError(f"unsupported order direction: {direction}") + order = OrderBy.asc(order) if normalized == "asc" else OrderBy.desc(order) + elif direction is not None: + raise TypeError("direction is only valid when order is a field name") self.order_by_items.append(order) return self diff --git a/src/teaql/core/value.py b/src/teaql/core/value.py index b2ce12c..7efbbe4 100644 --- a/src/teaql/core/value.py +++ b/src/teaql/core/value.py @@ -88,6 +88,8 @@ def Text(v: str) -> 'Value': def Json(v: Any) -> 'Value': return Value(v, DataType.Json) + JSON = Json + @staticmethod def Date(v: date) -> 'Value': return Value(v) @@ -96,6 +98,10 @@ def Date(v: date) -> 'Value': def Timestamp(v: Timestamp) -> 'Value': return Value(v) + @staticmethod + def DateTime(v: Any) -> 'Value': + return Value.from_any(v) + @staticmethod def Object(v: Dict[str, 'Value']) -> 'Value': return Value(v, DataType.Json) # Or Object specific diff --git a/src/teaql/data_service/__init__.py b/src/teaql/data_service/__init__.py index 8716ef8..3a032c3 100644 --- a/src/teaql/data_service/__init__.py +++ b/src/teaql/data_service/__init__.py @@ -165,6 +165,41 @@ class DataService(QueryExecutor, MutationExecutor, Protocol): pass +def _generated_schema_provider(): + from teaql.provider.sqlite import SimpleSchemaProvider + return SimpleSchemaProvider() + + +class SQLiteTeaQLClient: + """Stable high-level SQLite entry point used by generated workspaces.""" + def __new__(cls, database_url: str): + from teaql.provider.sqlite import create_sqlite_service + return create_sqlite_service(database_url, _generated_schema_provider()) + + +class TeaQLClient(SQLiteTeaQLClient): + """Portable local client backed by the packaged SQLite provider.""" + pass + + +class PostgreSQLTeaQLClient: + def __new__(cls, database_url: str): + from teaql.provider.postgres.dialect import PostgresDialect + from teaql.provider.postgres.transport import PostgresTransport + from teaql.sql.executor import SqlDataServiceExecutor + return SqlDataServiceExecutor(PostgresDialect(), PostgresTransport(database_url), + _generated_schema_provider()) + + +class MySQLTeaQLClient: + def __new__(cls, database_url: str): + from teaql.provider.mysql.dialect import MysqlDialect + from teaql.provider.mysql.transport import MysqlTransport + from teaql.sql.executor import SqlDataServiceExecutor + return SqlDataServiceExecutor(MysqlDialect(), MysqlTransport(database_url), + _generated_schema_provider()) + + __all__ = [ "DataServiceCapabilities", "QueryRequest", @@ -186,4 +221,5 @@ class DataService(QueryExecutor, MutationExecutor, Protocol): "TransactionExecutor", "IdGeneratorExecutor", "DataService" + , "TeaQLClient", "SQLiteTeaQLClient", "PostgreSQLTeaQLClient", "MySQLTeaQLClient" ] diff --git a/src/teaql/runtime/__init__.py b/src/teaql/runtime/__init__.py index 94264bb..e0c46fa 100644 --- a/src/teaql/runtime/__init__.py +++ b/src/teaql/runtime/__init__.py @@ -10,10 +10,11 @@ from .module import RuntimeModule, DefaultEntityDataServiceBehavior from .store import DataStore from .audit import RawAuditEvent, SafeAuditEvent, MutationAuditKind -from .i18n import CheckException, CheckResult, I18nCatalog, Locale, ObjectLocation, UnsupportedLocaleError +from .i18n import CheckException, CheckResult, I18nCatalog, JsonFieldNamingProfile, Locale, ObjectLocation, UnsupportedLocaleError +from .wire_fields import NormalizedWireInput, WireEntityMetadata, WireFieldMetadata, WireInputError, create_wire_entity_metadata, encode_wire_output, normalize_wire_input, retain_submitted_paths from teaql.core.entity import EntityKey, EntityChangeSet, EntityRoot -__all__ = ["EntityKey", "EntityChangeSet", "EntityRoot", "ContextEntityRef", "ContextRootError", "CheckException", "CheckResult", "I18nCatalog", "Locale", "ObjectLocation", "UnsupportedLocaleError", "UserContext", "TeaqlRuntime", "SqlLogEntry", "SqlLogOperation", "DiagnosticSqlLogSink", "TextDiagnosticSqlLogSink", "ServiceRuntimeFromEnv", "RuntimeModule", "DataStore", "RawAuditEvent", "SafeAuditEvent", "MutationAuditKind", "ContextTools", "ExecutableHttpTool", "HTTP_TOOL", "HttpIntentPhase", "HttpTool", "HttpToolProvider", "ToolDeniedError", "ToolError", "ToolPolicy", "ToolRisk", "Tools", "ToolToken", "ToolUnavailableError"] +__all__ = ["WireFieldMetadata", "WireEntityMetadata", "NormalizedWireInput", "WireInputError", "create_wire_entity_metadata", "normalize_wire_input", "encode_wire_output", "retain_submitted_paths", "EntityKey", "EntityChangeSet", "EntityRoot", "ContextEntityRef", "ContextRootError", "CheckException", "CheckResult", "I18nCatalog", "JsonFieldNamingProfile", "Locale", "ObjectLocation", "UnsupportedLocaleError", "UserContext", "TeaqlRuntime", "SqlLogEntry", "SqlLogOperation", "DiagnosticSqlLogSink", "TextDiagnosticSqlLogSink", "ServiceRuntimeFromEnv", "RuntimeModule", "DataStore", "RawAuditEvent", "SafeAuditEvent", "MutationAuditKind", "ContextTools", "ExecutableHttpTool", "HTTP_TOOL", "HttpIntentPhase", "HttpTool", "HttpToolProvider", "ToolDeniedError", "ToolError", "ToolPolicy", "ToolRisk", "Tools", "ToolToken", "ToolUnavailableError"] def __getattr__(name): diff --git a/src/teaql/runtime/context.py b/src/teaql/runtime/context.py index 626298d..1e5b8e8 100644 --- a/src/teaql/runtime/context.py +++ b/src/teaql/runtime/context.py @@ -104,6 +104,7 @@ def __init__(self): self._graph_rollback_actions: List[Any] = [] self._fix_evidence_current: List[FixEvidence] = [] self._fix_evidence_last: List[FixEvidence] = [] + self._checked_mutations = set() def begin_fix_evidence(self): self._fix_evidence_current = [] @@ -303,6 +304,8 @@ def with_metadata(self, metadata: Any) -> 'UserContext': def insert_resource(self, resource_type: str, resource: Any): self._resources[resource_type] = resource + if resource_type == "dataService" and hasattr(resource, "_ensure_schema"): + self._resources["schema_provider"] = resource return self def get_resource(self, resource_type: str) -> Optional[Any]: @@ -420,6 +423,20 @@ def with_custom_event_sink(self, sink: Any) -> 'UserContext': def set_custom_event_sink(self, sink: Any): self._app_audit_sink = sink + def initialize_audit(self, raw_sink: Any, app_sink: Any = None) -> 'UserContext': + """Compatibility entry point for generated applications.""" + self._set_standard_audit_sink(raw_sink) + self._app_audit_sink = app_sink + return self + + def configure_audit_policy(self, entity: str, mask_fields: List[str], + value_max_len: Optional[int] = None) -> 'UserContext': + descriptor = self.entity(entity) + if descriptor is not None: + descriptor.audit_mask_fields(mask_fields) + descriptor.audit_value_max_len(value_max_len) + return self + def with_internal_id_generator(self, gen: Any) -> 'UserContext': self.insert_resource("internal_id_generator", gen) return self @@ -577,6 +594,16 @@ def check_and_fix_mutation(self, mutation: Any): self.finish_fix_evidence() self._resources.pop("fix_operation", None) + def mark_mutation_checked(self, mutation: Any) -> None: + self._checked_mutations.add(id(mutation)) + + def consume_mutation_checked(self, mutation: Any) -> bool: + key = id(mutation) + if key not in self._checked_mutations: + return False + self._checked_mutations.remove(key) + return True + def translate_check_results(self, results: Any): for r in results: self.i18n_catalog().translate_check_result(r, self.language()) diff --git a/src/teaql/runtime/i18n.py b/src/teaql/runtime/i18n.py index 23029fb..1151359 100644 --- a/src/teaql/runtime/i18n.py +++ b/src/teaql/runtime/i18n.py @@ -6,6 +6,22 @@ from importlib.resources import files from typing import Any, Mapping +class JsonFieldNamingProfile(str, Enum): + CAMEL_CASE = "camelCase" + SNAKE_CASE = "snake_case" + PASCAL_CASE = "PascalCase" + + @classmethod + def parse(cls, value: str | None) -> JsonFieldNamingProfile: + if value is None or value == "": return cls.CAMEL_CASE + try: return cls(value) + except ValueError as error: raise ValueError(f"unsupported json_field_naming: {value}") from error + + def render(self, canonical_name: str) -> str: + if self is self.SNAKE_CASE: return canonical_name + camel = _lower_camel(canonical_name) + return camel if self is self.CAMEL_CASE or not camel else camel[:1].upper() + camel[1:] + @dataclass(frozen=True, eq=False) class ObjectLocation: """A checker location whose source of truth is the canonical KSML path.""" @@ -30,9 +46,12 @@ def native_path(self) -> str: @builtins.property def instance_path(self) -> str: + return self.instance_path_for(JsonFieldNamingProfile.CAMEL_CASE) + + def instance_path_for(self, profile: JsonFieldNamingProfile) -> str: parts = [] for kind, value in self.segments: - text = str(value) if kind == "index" else _lower_camel(str(value)) + text = str(value) if kind == "index" else profile.render(str(value)) parts.append(text.replace("~", "~0").replace("/", "~1")) return "" if not parts else "/" + "/".join(parts) @@ -94,10 +113,17 @@ def parse(cls, code): @dataclass class CheckResult: - rule_id:str; location:Any; input_value:Any=None; system_value:Any=None; message:str|None=None + rule_id:str; location:Any; input_value:Any=None; system_value:Any=None; message:str|None=None; entity_type:str|None=None; source_instance_path:str|None=None def __post_init__(self): if isinstance(self.location, str): self.location = ObjectLocation.from_model_path(self.location) + def to_wire(self, profile:JsonFieldNamingProfile=JsonFieldNamingProfile.CAMEL_CASE)->dict[str,Any]: + return {key:value for key,value in { + "ruleId":self.rule_id,"entityType":self.entity_type, + "location":[{"kind":kind, "name":value} if kind=="property" else {"kind":kind,"index":value} for kind,value in self.location.segments], + "instancePath":self.location.instance_path_for(profile),"sourceInstancePath":self.source_instance_path, + "inputValue":self.input_value,"systemValue":self.system_value,"message":self.message, + }.items() if value is not None} class CheckException(Exception): """Stable machine-readable model validation failure.""" diff --git a/src/teaql/runtime/module.py b/src/teaql/runtime/module.py index 241b7d6..14cbe9d 100644 --- a/src/teaql/runtime/module.py +++ b/src/teaql/runtime/module.py @@ -9,6 +9,8 @@ def __init__(self): self._audit_sinks: List[Any] = [] self._checkers: Dict[str, Any] = {} self._generated_bootstraps: List[Any] = [] + self._schema_entities: List[Any] = [] + self._wire_metadata: Dict[str, Any] = {} @classmethod def new(cls) -> 'RuntimeModule': @@ -54,6 +56,15 @@ def generated_bootstrap(self, bootstrap: Any) -> 'RuntimeModule': self._generated_bootstraps.append(bootstrap) return self + def wire_metadata(self, entity: str, metadata: Any) -> 'RuntimeModule': + self._wire_metadata[entity] = metadata + return self + + def schema_entity(self, descriptor: Any) -> 'RuntimeModule': + """Register generated storage metadata without embedding a provider.""" + self._schema_entities.append(descriptor) + return self + def and_module(self, other: 'RuntimeModule') -> 'RuntimeModule': combined = RuntimeModule.new() combined._entities = [*self._entities, *other._entities] @@ -65,6 +76,8 @@ def and_module(self, other: 'RuntimeModule') -> 'RuntimeModule': *self._generated_bootstraps, *other._generated_bootstraps, ] + combined._schema_entities = [*self._schema_entities, *other._schema_entities] + combined._wire_metadata = {**self._wire_metadata, **other._wire_metadata} combined._initial_graphs = [ *getattr(self, '_initial_graphs', []), *getattr(other, '_initial_graphs', []), @@ -78,7 +91,8 @@ def and_module(self, other: 'RuntimeModule') -> 'RuntimeModule': def apply_to(self, context: UserContext): for name, dep in self._dependencies.items(): context.insert_resource(name, dep) - for entity in self._entities: + installed_entities = self._schema_entities or self._entities + for entity in installed_entities: context.register_entity(entity) if hasattr(self, '_initial_graphs'): for graph in self._initial_graphs: @@ -86,8 +100,10 @@ def apply_to(self, context: UserContext): if hasattr(self, '_root_graphs'): for graph in self._root_graphs: context.add_root_graph(graph) - context.insert_resource("entities", self._entities) + context.insert_resource("entities", installed_entities) + context.insert_resource("entity_classes", self._entities) context.insert_resource("behaviors", self._behaviors) + context.insert_resource("wireMetadata", dict(self._wire_metadata)) context.insert_resource("_teaql_generated_bootstraps", tuple(self._generated_bootstraps)) if self._checkers: checkers = dict(self._checkers) diff --git a/src/teaql/runtime/wire_fields.py b/src/teaql/runtime/wire_fields.py new file mode 100644 index 0000000..5cdf4d1 --- /dev/null +++ b/src/teaql/runtime/wire_fields.py @@ -0,0 +1,66 @@ +from __future__ import annotations +from dataclasses import dataclass +from types import MappingProxyType +from .i18n import CheckResult, JsonFieldNamingProfile + +@dataclass(frozen=True) +class WireFieldMetadata: + canonical_name: str + wire_name: str + aliases: tuple[str, ...] = () + +@dataclass(frozen=True) +class WireEntityMetadata: + entity_type: str + profile: JsonFieldNamingProfile + fields: object + +@dataclass(frozen=True) +class NormalizedWireInput: + values: object + source_instance_paths: object + +class WireInputError(ValueError): + def __init__(self, code, instance_path, message): + super().__init__(message); self.code = code; self.instance_path = instance_path + +def create_wire_entity_metadata(entity_type, canonical_fields, profile=JsonFieldNamingProfile.CAMEL_CASE, aliases=None): + aliases = aliases or {}; fields = {}; spellings = {} + for canonical_name in canonical_fields: + field = WireFieldMetadata(canonical_name, profile.render(canonical_name), tuple(aliases.get(canonical_name, ()))) + for spelling in (field.wire_name, *field.aliases): + previous = spellings.get(spelling) + if previous is not None and previous != canonical_name: + raise ValueError(f"Wire field spelling '{spelling}' maps to both '{previous}' and '{canonical_name}'") + spellings[spelling] = canonical_name + fields[canonical_name] = field + return WireEntityMetadata(entity_type, profile, MappingProxyType(fields)) + +def normalize_wire_input(input_value, metadata, parent_pointer=""): + lookup = {spelling: field for field in metadata.fields.values() for spelling in (field.wire_name, *field.aliases)} + values = {}; paths = {}; submitted = {} + for submitted_name, value in input_value.items(): + pointer = f"{parent_pointer}/{_escape_pointer(submitted_name)}"; field = lookup.get(submitted_name) + if field is None: raise WireInputError("WIRE_UNKNOWN_FIELD", pointer, f"Unknown {metadata.entity_type} field '{submitted_name}'") + previous = submitted.get(field.canonical_name) + if previous is not None: raise WireInputError("WIRE_FIELD_COLLISION", pointer, f"Fields '{previous}' and '{submitted_name}' both map to canonical field '{field.canonical_name}'") + submitted[field.canonical_name] = submitted_name; values[field.canonical_name] = value + if submitted_name != field.wire_name: paths[field.canonical_name] = pointer + return NormalizedWireInput(MappingProxyType(values), MappingProxyType(paths)) + +def encode_wire_output(values, metadata): + output = {} + for canonical_name, value in values.items(): + field = metadata.fields.get(canonical_name) + if field is None: raise ValueError(f"Unknown canonical {metadata.entity_type} field '{canonical_name}'") + output[field.wire_name] = value + return MappingProxyType(output) + +def retain_submitted_paths(results, normalized): + retained = [] + for result in results: + canonical = next((value for kind, value in result.location.segments if kind == "property"), None) + retained.append(CheckResult(result.rule_id, result.location, result.input_value, result.system_value, result.message, result.entity_type, normalized.source_instance_paths.get(canonical, result.source_instance_path))) + return retained + +def _escape_pointer(value): return value.replace("~", "~0").replace("/", "~1") diff --git a/src/teaql/sql/dialect.py b/src/teaql/sql/dialect.py index d30bf33..e70ae5e 100644 --- a/src/teaql/sql/dialect.py +++ b/src/teaql/sql/dialect.py @@ -253,13 +253,13 @@ def compile_update(self, entity: EntityDescriptor, command: UpdateCommand) -> Co return CompiledQuery(sql=sql, params=params) def compile_delete(self, entity: EntityDescriptor, command: DeleteCommand) -> CompiledQuery: - id_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_id_val', False) or (callable(getattr(p, 'is_id', None)) and p.is_id())), None) + id_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_id_val', False) or getattr(p, '_is_id', False)), None) if not id_property: raise MissingIdPropertyError(entity._name) params = [] table_name = getattr(entity, 'table_name_val', entity._name) - version_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_version_val', False) or (callable(getattr(p, 'is_version', None)) and p.is_version())), None) + version_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_version_val', False) or getattr(p, '_is_version', False)), None) if command.soft_delete: if not version_property: @@ -295,11 +295,11 @@ def compile_recover(self, entity: EntityDescriptor, command: RecoverCommand) -> if command.expected_version_val >= 0: raise InvalidRecoverVersionError(command.expected_version_val) - id_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_id_val', False) or (callable(getattr(p, 'is_id', None)) and p.is_id())), None) + id_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_id_val', False) or getattr(p, '_is_id', False)), None) if not id_property: raise MissingIdPropertyError(entity._name) - version_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_version_val', False) or (callable(getattr(p, 'is_version', None)) and p.is_version())), None) + version_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, 'is_version_val', False) or getattr(p, '_is_version', False)), None) if not version_property: raise MissingVersionPropertyError(entity._name) diff --git a/src/teaql/sql/executor.py b/src/teaql/sql/executor.py index 00aae48..aeb07a6 100644 --- a/src/teaql/sql/executor.py +++ b/src/teaql/sql/executor.py @@ -104,6 +104,19 @@ def __init__(self, dialect: SqlDialect, transport: SqlTransport, schema_provider self.transport = transport self.schema_provider = schema_provider + def _sync_generated_schema(self, context: 'UserContext') -> None: + register = getattr(self.schema_provider, "register_entity", None) + if callable(register) and context is not None: + for entity in context.all_entities(): + register(entity) + + async def close(self) -> None: + close = getattr(self.transport, "close", None) + if callable(close): + result = close() + if hasattr(result, "__await__"): + await result + def _resolve_subquery_entities(self, expr) -> None: if expr is None: return @@ -183,6 +196,7 @@ async def query(self, context: 'UserContext', request: QueryRequest) -> QueryRes ) async def _query(self, context: 'UserContext', request: QueryRequest) -> QueryResult: + self._sync_generated_schema(context) request.query.prepare_for_list() execution_query, retained_order, retained_empty = await self._prepare_id_set_page(context, request.query) if retained_empty: @@ -565,6 +579,7 @@ def _attach_empty_relation_aggregate(self, parents, aggregate, query): parent[aggregate.alias] = value async def mutate(self, context: 'UserContext', request: MutationRequest) -> MutationResult: + self._sync_generated_schema(context) entity = getattr(request._data, "entity", "unknown") kind = type(request._data).__name__.replace("Command", "").lower() if context is not None: @@ -618,6 +633,12 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut else: await self.ensure_id_floor( req_data.entity, int(req_data.values[id_prop.name].val)) + version_prop = next(( + prop for prop in getattr(entity_desc, "properties", []) + if getattr(prop, "_is_version", False) + ), None) + if version_prop is not None and version_prop.name not in req_data.values: + req_data.values[version_prop.name] = Value.from_any(1) try: if isinstance(req_data, InsertCommand): @@ -722,8 +743,10 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut parameters=list(compiled.params), affected_rows=affected_rows, trace_chain=trace_path, - comment=request.comment(), - audit_reason=request.comment(), + comment=(request.comment() if callable(getattr(request, "comment", None)) + else getattr(request, "comment", None)), + audit_reason=(request.comment() if callable(getattr(request, "comment", None)) + else getattr(request, "comment", None)), debug_query=compiled.debug_sql(self.dialect.kind()) ) if context is not None: @@ -856,6 +879,7 @@ async def _ensure_schema(self, context: 'UserContext', capability: object) -> No from teaql.runtime._schema_capability import SCHEMA_CAPABILITY if capability is not SCHEMA_CAPABILITY: raise PermissionError("Ensure Schema is available only through UserContext.ensure_schema()") + self._sync_generated_schema(context) enable_soundex = getattr(self.transport, "enable_soundex", None) if callable(enable_soundex): await enable_soundex() diff --git a/tests/core/test_query.py b/tests/core/test_query.py index 2cc3fd0..a25c5d6 100644 --- a/tests/core/test_query.py +++ b/tests/core/test_query.py @@ -15,6 +15,11 @@ def test_order_by_builder(): assert ob_desc.field_name == "created_at" assert ob_desc.direction == SortDirection.Desc + query = SelectQuery.new("User").order_by("created_at", "desc") + assert query.order_by_items == [ob_desc] + with pytest.raises(ValueError, match="unsupported order direction"): + query.order_by("id", "sideways") + def test_id_set_pagination_is_explicit_local_and_validated(): query = SelectQuery.new("Order") assert query.id_set_pagination is None diff --git a/tests/runtime/test_i18n.py b/tests/runtime/test_i18n.py index b7a825d..4ccea66 100644 --- a/tests/runtime/test_i18n.py +++ b/tests/runtime/test_i18n.py @@ -2,16 +2,26 @@ from importlib.resources import files import pytest -from teaql.runtime import CheckResult, I18nCatalog, Locale, ObjectLocation, UnsupportedLocaleError, UserContext +from teaql.runtime import CheckResult, I18nCatalog, JsonFieldNamingProfile, Locale, ObjectLocation, UnsupportedLocaleError, UserContext def test_object_location_preserves_all_three_path_dialects(): location = ObjectLocation().property("order_items").index(2).property("user_url") assert location.model_path == "order_items[2].user_url" assert location.native_path == "order_items[2].user_url" assert location.instance_path == "/orderItems/2/userUrl" + assert location.instance_path_for(JsonFieldNamingProfile.SNAKE_CASE) == "/order_items/2/user_url" + assert location.instance_path_for(JsonFieldNamingProfile.PASCAL_CASE) == "/OrderItems/2/UserUrl" escaped = ObjectLocation().property("a/b~c") assert escaped.instance_path == "/a~1b~0c" +def test_wire_checker_result_preserves_submitted_alias(): + result = CheckResult("required", ObjectLocation().property("user_url"), + entity_type="customer_account", source_instance_path="/user_url") + wire = result.to_wire() + assert wire["instancePath"] == "/userUrl" + assert wire["sourceInstancePath"] == "/user_url" + assert wire["location"] == [{"kind":"property", "name":"user_url"}] + def test_fifteen_locales_times_five_checker_rules(): catalog = I18nCatalog.builtin() results = [CheckResult("required", "name"), CheckResult("min", "age", 1, 2), diff --git a/tests/runtime/test_wire_fields.py b/tests/runtime/test_wire_fields.py new file mode 100644 index 0000000..54f00e6 --- /dev/null +++ b/tests/runtime/test_wire_fields.py @@ -0,0 +1,21 @@ +import pytest +from teaql.runtime import CheckResult, JsonFieldNamingProfile, ObjectLocation, WireInputError, create_wire_entity_metadata, encode_wire_output, normalize_wire_input, retain_submitted_paths + +def metadata(): + return create_wire_entity_metadata("School", ["name", "school_type"], JsonFieldNamingProfile.CAMEL_CASE, {"school_type": ["school_type"]}) + +def test_normalizes_declared_alias_and_retains_provenance(): + normalized = normalize_wire_input({"name": "Riverside", "school_type": 1001}, metadata(), "/school") + assert dict(normalized.values) == {"name": "Riverside", "school_type": 1001} + result = CheckResult("required", ObjectLocation().property("school_type")) + assert retain_submitted_paths([result], normalized)[0].source_instance_path == "/school/school_type" + +def test_rejects_unknown_and_collision(): + with pytest.raises(WireInputError) as unknown: normalize_wire_input({"not/known": 1}, metadata()) + assert (unknown.value.code, unknown.value.instance_path) == ("WIRE_UNKNOWN_FIELD", "/not~1known") + with pytest.raises(WireInputError) as collision: normalize_wire_input({"schoolType": 1001, "school_type": 1002}, metadata()) + assert collision.value.code == "WIRE_FIELD_COLLISION" + +def test_encodes_wire_output(): + assert dict(encode_wire_output({"school_type": 1001}, metadata())) == {"schoolType": 1001} + with pytest.raises(ValueError, match="Unknown canonical"): encode_wire_output({"missing": 1}, metadata())