From d2f085cfe3027482dc1c5d5b3c3a95a425c04fa9 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 08:47:32 +0800 Subject: [PATCH 01/13] feat(python): render expanded SQL with field-aware masking --- .gitattributes | 1 + examples/conformance/app/main.py | 5 +- examples/conformance/app/masking_lifecycle.py | 148 ++++ examples/conformance/models/platform.py | 2 +- examples/conformance/models/work_item.py | 1 - examples/conformance/pyproject.toml | 2 +- .../conformance/requests/platform_request.py | 36 +- .../conformance/requests/work_item_request.py | 36 +- examples/conformance/runtime_module.py | 6 +- examples/conformance/test_sql_log_intent.py | 33 + examples/order-management/model.xml | 2 +- .../models/commerce_platform.py | 2 +- .../python-lib-core/models/customer.py | 2 +- .../python-lib-core/models/customer_order.py | 4 +- .../python-lib-core/models/order_line.py | 1 - .../models/order_search_preset.py | 1 - .../python-lib-core/models/order_status.py | 2 +- .../python-lib-core/models/product.py | 2 +- .../python-lib-core/pyproject.toml | 2 +- .../requests/commerce_platform_request.py | 224 ++++-- .../requests/customer_order_request.py | 126 +-- .../requests/customer_request.py | 96 ++- .../requests/order_line_request.py | 68 +- .../requests/order_search_preset_request.py | 44 +- .../requests/order_status_request.py | 112 +-- .../requests/product_request.py | 98 ++- .../python-lib-core/runtime_module.py | 21 +- .../order-management/test_sql_log_intent.py | 37 + examples/school-management/app/main.py | 3 +- examples/school-management/models/platform.py | 2 +- examples/school-management/models/school.py | 2 +- .../school-management/models/school_type.py | 2 +- examples/school-management/pyproject.toml | 2 +- .../requests/platform_request.py | 159 ++-- .../requests/school_request.py | 81 +- .../requests/school_type_request.py | 123 +-- examples/school-management/runtime_module.py | 9 +- .../school-management/test_sql_log_intent.py | 32 + examples/task_board/generated/E.py | 217 ++++++ examples/task_board/generated/Q.py | 38 + examples/task_board/generated/model.xml | 55 ++ .../task_board/generated/models/platform.py | 307 +++++++- examples/task_board/generated/models/task.py | 306 +++++++- .../generated/models/task_execution_log.py | 250 +++++- .../generated/models/task_status.py | 341 +++++++- examples/task_board/generated/pyproject.toml | 16 + .../generated/requests/platform_request.py | 726 +++++++++++++++++- .../requests/task_execution_log_request.py | 455 ++++++++++- .../generated/requests/task_request.py | 465 ++++++++++- .../generated/requests/task_status_request.py | 714 ++++++++++++++++- .../task_board/generated/runtime_module.py | 368 +++++++++ examples/task_board/generated/teaql-i18n.json | 123 +++ examples/task_board/main.py | 201 ++--- examples/task_board/test_task_board.py | 47 ++ scripts/verify-examples.sh | 4 + src/teaql/core/meta.py | 7 + src/teaql/data_service/__init__.py | 5 + src/teaql/provider/postgres/dialect.py | 2 +- src/teaql/runtime/audit.py | 13 +- src/teaql/runtime/context.py | 27 +- src/teaql/runtime/log_privacy.py | 149 +++- src/teaql/sql/dialect.py | 78 +- src/teaql/sql/executor.py | 261 ++++--- src/teaql/sql/types.py | 67 +- test-vectors/masking-v1.tsv | 11 + tests/core/test_context.py | 4 +- tests/runtime/test_log_privacy.py | 6 +- tests/runtime/test_masking_contract.py | 66 ++ tests/runtime/test_relation_masking.py | 150 ++++ tests/runtime/test_sql_mask_lifecycle.py | 421 ++++++++++ tests/runtime/test_sql_masking_policy.py | 223 ++++++ 71 files changed, 6595 insertions(+), 1057 deletions(-) create mode 100644 .gitattributes create mode 100644 examples/conformance/app/masking_lifecycle.py create mode 100644 examples/conformance/test_sql_log_intent.py create mode 100644 examples/order-management/test_sql_log_intent.py create mode 100644 examples/school-management/test_sql_log_intent.py create mode 100644 examples/task_board/generated/E.py create mode 100644 examples/task_board/generated/Q.py create mode 100644 examples/task_board/generated/model.xml create mode 100644 examples/task_board/generated/pyproject.toml create mode 100644 examples/task_board/generated/runtime_module.py create mode 100644 examples/task_board/generated/teaql-i18n.json create mode 100644 examples/task_board/test_task_board.py create mode 100644 test-vectors/masking-v1.tsv create mode 100644 tests/runtime/test_masking_contract.py create mode 100644 tests/runtime/test_relation_masking.py create mode 100644 tests/runtime/test_sql_mask_lifecycle.py create mode 100644 tests/runtime/test_sql_masking_policy.py diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 0000000..d4a292f --- /dev/null +++ b/.gitattributes @@ -0,0 +1 @@ +test-vectors/*.tsv whitespace=-blank-at-eol diff --git a/examples/conformance/app/main.py b/examples/conformance/app/main.py index df49be5..3d43742 100644 --- a/examples/conformance/app/main.py +++ b/examples/conformance/app/main.py @@ -1,4 +1,5 @@ import asyncio +import os from pathlib import Path import sys @@ -12,9 +13,11 @@ from teaql.data_service import SQLiteTeaQLClient from teaql.core import EntityKey, EntityRoot from teaql.runtime import CheckException, UserContext +from app.masking_lifecycle import verify_masking_lifecycle async def main() -> None: + await verify_masking_lifecycle() order_key = EntityKey("Order", 1) execution_key = EntityKey("InferenceExecution", 1) target_ledger = EntityRoot() @@ -27,7 +30,7 @@ async def main() -> None: assert target_ledger.original_version(execution_key) == 9 print("PASS Mutation ledger identity (same ID, different entity types keep versions 3/9)") - database = ROOT / ".local" / "conformance.sqlite" + database = Path(os.environ.get("TEAQL_CONFORMANCE_DB", ROOT / ".local" / "conformance.sqlite")) database.parent.mkdir(parents=True, exist_ok=True) database.unlink(missing_ok=True) client = SQLiteTeaQLClient(str(database)) diff --git a/examples/conformance/app/masking_lifecycle.py b/examples/conformance/app/masking_lifecycle.py new file mode 100644 index 0000000..1b15783 --- /dev/null +++ b/examples/conformance/app/masking_lifecycle.py @@ -0,0 +1,148 @@ +"""Runtime-owned fixture; does not modify generated domain-library code.""" +from contextlib import aclosing +from pathlib import Path +from tempfile import TemporaryDirectory +from types import SimpleNamespace + +from teaql.core.expr import Expr +from teaql.core.meta import EntityDescriptor, PropertyDescriptor, RelationDescriptor +from teaql.core.mutation import InsertCommand, MutationRequest, TraceNode +from teaql.core.query import SelectQuery +from teaql.core.value import DataType +from teaql.data_service import QueryRequest +from teaql.provider.sqlite import SimpleSchemaProvider, create_sqlite_service +from teaql.runtime import RuntimeModule +from teaql.runtime.context import TextDiagnosticSqlLogSink +from teaql.sql.executor import TransportError +from teaql.sql.types import CompiledQuery + + +async def verify_masking_lifecycle(): + with TemporaryDirectory(prefix='teaql-mask-lifecycle-') as directory: + entity = (EntityDescriptor('MaskCustomer').table_name('mask_customer_data') + .property(PropertyDescriptor('id', DataType.I64).is_id()) + .property(PropertyDescriptor('version', DataType.I64).is_version()) + .property(PropertyDescriptor('display_name', DataType.Text)) + .audit_mask_fields(['display_name'])) + entity.relation(RelationDescriptor('children', 'MaskChild').foreign('parent_id').many()) + child = (EntityDescriptor('MaskChild').table_name('mask_child_data') + .property(PropertyDescriptor('id', DataType.I64).is_id()) + .property(PropertyDescriptor('version', DataType.I64).is_version()) + .property(PropertyDescriptor('parent_id', DataType.I64))) + provider = SimpleSchemaProvider() + provider.register_entity(entity) + provider.register_entity(child) + service = create_sqlite_service(str(Path(directory) / 'mask.sqlite'), provider) + context = RuntimeModule.new().entity(entity).entity(child).into_context().with_schema_provider(service) + await context.ensure_schema() + output, entries = [], [] + log_file = Path(directory) / 'sql.log' + def write_line(text): + output.append(text) + with log_file.open('a', encoding='utf-8') as destination: + destination.write(text + '\n') + sink = TextDiagnosticSqlLogSink(write_line) + def capture(entry): + entries.append(entry) + sink.write(entry) + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=capture)) + async def insert(entity_id, target=None): + command = (InsertCommand('MaskCustomer').value('id', entity_id).value('version', 1) + .value('display_name', 'Riverside')) + command.trace_chain = [TraceNode(comment='what: seed Riverside lifecycle fixture')] + return await (target or service).mutate(context, MutationRequest(command)) + try: + for entity_id in [1, 2, 3]: + await insert(entity_id) + entries.clear() + output.clear() + request = QueryRequest(SelectQuery('MaskCustomer') + .filter(Expr.eq('display_name', 'Riverside')).limit(3) + ).comment('what: inspect customers').purpose('why: verify stream lifecycle') + # async-for break alone does not promise immediate generator close + # in Python. aclosing makes ownership explicit and deterministic. + async with aclosing(service.query_stream(context, request, 1)) as stream: + async for chunk in stream: + assert chunk.rows[0]['display_name'] == 'Riverside' + break + assert len(entries) == 1 and entries[0].execution_outcome == 'cancelled' + assert entries[0].result_count == 1 + try: + await insert(1) + raise AssertionError('duplicate key unexpectedly succeeded') + except TransportError: + pass + assert entries[-1].execution_outcome == 'failure' + assert entries[-1].affected_rows is None + result = await service.query(context, request) + assert len(result.rows) == 3 + assert all(row['display_name'] == 'Riverside' for row in result.rows) + assert entries[-1].execution_outcome == 'success' + text = '\n'.join(output) + assert 'Riverside' not in text and 'Ri*****de' in text + assert 'what: inspect customers' in text and 'why: verify stream lifecycle' in text + + await service.transport.execute_sql(CompiledQuery( + 'CREATE TRIGGER remove_mask_probe AFTER INSERT ON mask_customer_data ' + 'WHEN NEW.id = 777 BEGIN DELETE FROM mask_customer_data WHERE id = NEW.id; END', [])) + entries.clear() + output.clear() + try: + await insert(777) + raise AssertionError('missing snapshot unexpectedly succeeded') + except TransportError: + pass + assert len(entries) == 2 + assert entries[0].execution_outcome == 'success' and entries[0].affected_rows == 1 + assert entries[1].execution_outcome == 'success' and entries[1].result_count == 0 + assert 'what: seed' in entries[1].audit_reason + assert 'Riverside' not in repr(entries) and 'Riverside' not in '\n'.join(output) + + context.insert_resource('dataService', service) + entries.clear() + async def partial_graph(): + target = context.require_resource('dataService') + await insert(30, target) + await insert(777, target) + await insert(31, target) + try: + await context.execute_graph_save(partial_graph) + raise AssertionError('partial graph unexpectedly committed') + except TransportError: + pass + assert [entry.execution_outcome for entry in entries] == ['success','success','success'] + assert entries[-1].result_count == 0 + assert not await service.transport.fetch_all_sql(CompiledQuery( + 'SELECT id FROM mask_customer_data WHERE id IN (30,31,777)', [])) + assert 'Riverside' not in repr(entries) + await insert(32) + assert len(await service.transport.fetch_all_sql(CompiledQuery('SELECT id FROM mask_customer_data', []))) == 4 + + command = (InsertCommand('MaskChild').value('id', 1).value('version', 1).value('parent_id', 1)) + command.trace_chain = [TraceNode(comment='what: seed child for relation verification')] + await service.mutate(context, MutationRequest(command)) + graph_request = QueryRequest(SelectQuery('MaskCustomer').project('id') + .filter(Expr.eq('id', 1)).and_filter(Expr.eq('display_name', 'Riverside')) + .relation_query('children', SelectQuery('MaskChild').project('id').limit(2)).limit(1) + ).comment('what: load Riverside graph').purpose('why: verify inherited relation masking') + entries.clear() + await service.transport.execute_sql(CompiledQuery( + 'ALTER TABLE mask_child_data RENAME TO mask_child_unavailable', [])) + try: + try: + await service.query(context, graph_request) + raise AssertionError('missing relation table unexpectedly succeeded') + except TransportError: + pass + finally: + await service.transport.execute_sql(CompiledQuery( + 'ALTER TABLE mask_child_unavailable RENAME TO mask_child_data', [])) + assert [entry.execution_outcome for entry in entries] == ['success', 'failure'] + assert 'what: load' in entries[-1].comment + assert 'mask_child_data' in entries[-1].debug_sql + assert 'Riverside' not in repr(entries) + log_file.read_text(encoding='utf-8') + restored = await service.query(context, graph_request) + assert restored.rows[0]['children'][0]['id'] == 1 + finally: + await service.close() + print('PASS Python masked SQL lifecycle: stream, readback intent, SQLite failure, partial graph rollback, relation file/custom sinks and reuse') diff --git a/examples/conformance/models/platform.py b/examples/conformance/models/platform.py index 33aaa40..24f3e1b 100644 --- a/examples/conformance/models/platform.py +++ b/examples/conformance/models/platform.py @@ -246,4 +246,4 @@ def update_version(self, value): return self def work_item_list(self) -> list: self._loaded_fields.add("work_item_list") - return self._work_item_list \ No newline at end of file + return self._work_item_list diff --git a/examples/conformance/models/work_item.py b/examples/conformance/models/work_item.py index d3532c6..2ffd6b3 100644 --- a/examples/conformance/models/work_item.py +++ b/examples/conformance/models/work_item.py @@ -258,4 +258,3 @@ def update_platform(self, value): self._loaded_fields.add("platform") self._entity_root.set(self._teaql_entity_key(), "platform", Value.from_any(self.platform)) return self - diff --git a/examples/conformance/pyproject.toml b/examples/conformance/pyproject.toml index 182c898..565edd8 100644 --- a/examples/conformance/pyproject.toml +++ b/examples/conformance/pyproject.toml @@ -2,7 +2,7 @@ name = "runtime-example-conformance-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.7", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] diff --git a/examples/conformance/requests/platform_request.py b/examples/conformance/requests/platform_request.py index 6ca3eaf..2609c6e 100644 --- a/examples/conformance/requests/platform_request.py +++ b/examples/conformance/requests/platform_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -290,21 +286,21 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_work_item_list(self): from requests.work_item_request import WorkItemRequest @@ -322,13 +318,13 @@ def have_no_work_items(self): return self.without_work_item_list_matching(WorkItemRequest()) def with_work_item_list_matching(self, child_request): + child_request.query.projection = ["platform"] self.query.and_filter(in_subquery(column("id"), "WorkItem", child_request.query)) - child_request.query._projection = ["platform"] return self def without_work_item_list_matching(self, child_request): + child_request.query.projection = ["platform"] self.query.and_filter(not_in_subquery(column("id"), "WorkItem", child_request.query)) - child_request.query._projection = ["platform"] return self def count_work_items(self): return self.count_work_items_as("count_work_items") @@ -339,7 +335,9 @@ def count_work_items_as(self, alias: str): def count_work_items_with(self, alias: str, child_request): child_request.query.count_field("id", alias) - self.query.relation_aggregate("work_item_list", alias, child_request.query, True) + self.query.relation_aggregates.append( + RelationAggregate("work_item_list", alias, child_request.query, True) + ) return self @@ -366,7 +364,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -388,21 +386,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) query_root = EntityRoot() data = SmartList(Platform(_entity_root=query_root, **row) for row in row_result.rows) return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) @@ -421,6 +419,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): for row in chunk.rows: yield Platform(_entity_root=query_root, **row) diff --git a/examples/conformance/requests/work_item_request.py b/examples/conformance/requests/work_item_request.py index e6355c3..d3ed1ad 100644 --- a/examples/conformance/requests/work_item_request.py +++ b/examples/conformance/requests/work_item_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -102,12 +98,12 @@ def select_platform_with(self, child_request): self.query.relation_query("platform", child_request.query) return self def with_platform_matching(self, child_request): - child_request.query._projection = ["id"] + child_request.query.projection = ["id"] self.query.and_filter(in_subquery(column("platform"), "Platform", child_request.query)) return self def without_platform_matching(self, child_request): - child_request.query._projection = ["id"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_request.query)) return self @@ -400,35 +396,35 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_title(self): self.query.group_by("title") return self def group_by_title_as(self, ret_name: str): - self.query.group_by("title") + self.query.group_by("title") return self def group_by_description(self): self.query.group_by("description") return self def group_by_description_as(self, ret_name: str): - self.query.group_by("description") + self.query.group_by("description") return self def group_by_platform(self): self.query.group_by("platform") return self def group_by_platform_as(self, ret_name: str): - self.query.group_by("platform") + self.query.group_by("platform") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def facet_by_platform_as(self, name: str, request: QuerySelection, include_all_facets: bool = True): @@ -458,7 +454,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -480,21 +476,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) query_root = EntityRoot() data = SmartList(WorkItem(_entity_root=query_root, **row) for row in row_result.rows) return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) @@ -513,6 +509,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): for row in chunk.rows: yield WorkItem(_entity_root=query_root, **row) diff --git a/examples/conformance/runtime_module.py b/examples/conformance/runtime_module.py index 0ccf55c..247aba6 100644 --- a/examples/conformance/runtime_module.py +++ b/examples/conformance/runtime_module.py @@ -62,11 +62,13 @@ 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()) + .audit_mask_fields([]) + .table_name("platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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")) + .audit_mask_fields([]) + .table_name("work_item_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("title", DataType.Text).column_name("title").log_policy("plain").required()).property(PropertyDescriptor("description", DataType.Text).column_name("description").log_policy("plain")).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")) ) async def _ensure_generated_bootstrap_once(context): diff --git a/examples/conformance/test_sql_log_intent.py b/examples/conformance/test_sql_log_intent.py new file mode 100644 index 0000000..f8ffd94 --- /dev/null +++ b/examples/conformance/test_sql_log_intent.py @@ -0,0 +1,33 @@ +"""Verify generated query intent survives execution in the conformance app.""" +import os +from pathlib import Path +import subprocess +import sys +import tempfile +import unittest + + +class ConformanceSqlLogIntentTest(unittest.TestCase): + def test_generated_queries_retain_intent_at_sql_sink(self): + repo = Path(__file__).resolve().parents[2] + env = dict(os.environ) + env.pop("TEAQL_ALLOW_SENSITIVE_PLAINTEXT_LOGS", None) + env["PYTHONPATH"] = os.pathsep.join((str(repo / "examples" / "conformance"), str(repo / "src"))) + with tempfile.TemporaryDirectory(prefix="teaql-conformance-log-test-") as directory: + env["TEAQL_CONFORMANCE_DB"] = str(Path(directory) / "conformance.sqlite") + result = subprocess.run([sys.executable, "-m", "app.main"], cwd=repo, + env=env, capture_output=True, text=True, timeout=90) + output = result.stdout + result.stderr + self.assertEqual(0, result.returncode, output) + generated_flow = output.split("PASS Mutation ledger identity", 1)[-1] + queries = [line for line in generated_flow.splitlines() + if line.startswith("[TeaQL SQL]") and "[select]" in line] + self.assertGreater(len(queries), 0, output) + for line in queries: + self.assertNotIn("comment=None", line) + self.assertNotIn("purpose=None", line) + self.assertIn("PASS Python minimum runtime conformance: 8/8", output) + + +if __name__ == "__main__": + unittest.main() diff --git a/examples/order-management/model.xml b/examples/order-management/model.xml index 3003592..f0b3e2d 100644 --- a/examples/order-management/model.xml +++ b/examples/order-management/model.xml @@ -1,7 +1,7 @@ - + <_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-lib-core/models/commerce_platform.py b/examples/order-management/python-lib-core/models/commerce_platform.py index 95b64a0..a0b2a83 100644 --- a/examples/order-management/python-lib-core/models/commerce_platform.py +++ b/examples/order-management/python-lib-core/models/commerce_platform.py @@ -434,4 +434,4 @@ def order_line_list(self) -> 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 + return self._order_search_preset_list diff --git a/examples/order-management/python-lib-core/models/customer.py b/examples/order-management/python-lib-core/models/customer.py index 5029f09..c8cd7ab 100644 --- a/examples/order-management/python-lib-core/models/customer.py +++ b/examples/order-management/python-lib-core/models/customer.py @@ -325,4 +325,4 @@ def update_commerce_platform(self, value): def customer_order_list(self) -> list: self._loaded_fields.add("customer_order_list") - return self._customer_order_list \ No newline at end of file + return self._customer_order_list 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 5292dbc..7eb8c5c 100644 --- a/examples/order-management/python-lib-core/models/customer_order.py +++ b/examples/order-management/python-lib-core/models/customer_order.py @@ -376,10 +376,12 @@ def update_status(self, value): def update_status_to_pending(self): self.status = 1001 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_confirmed(self): self.status = 1002 self._loaded_fields.add("status") + self._entity_root.set(self._teaql_entity_key(), "status", Value.from_any(self.status)) return self @@ -398,4 +400,4 @@ def update_commerce_platform(self, value): def order_line_list(self) -> list: self._loaded_fields.add("order_line_list") - return self._order_line_list \ No newline at end of file + return self._order_line_list 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 fe06410..d754e9a 100644 --- a/examples/order-management/python-lib-core/models/order_line.py +++ b/examples/order-management/python-lib-core/models/order_line.py @@ -342,4 +342,3 @@ def update_commerce_platform(self, value): 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 4e8f370..7b774cc 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 @@ -334,4 +334,3 @@ def update_commerce_platform(self, value): 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 d71a584..d0eb8df 100644 --- a/examples/order-management/python-lib-core/models/order_status.py +++ b/examples/order-management/python-lib-core/models/order_status.py @@ -325,4 +325,4 @@ def update_commerce_platform(self, value): def customer_order_list(self) -> list: self._loaded_fields.add("customer_order_list") - return self._customer_order_list \ No newline at end of file + return self._customer_order_list diff --git a/examples/order-management/python-lib-core/models/product.py b/examples/order-management/python-lib-core/models/product.py index d82518a..6718e79 100644 --- a/examples/order-management/python-lib-core/models/product.py +++ b/examples/order-management/python-lib-core/models/product.py @@ -344,4 +344,4 @@ def update_commerce_platform(self, value): def order_line_list(self) -> list: self._loaded_fields.add("order_line_list") - return self._order_line_list \ No newline at end of file + return self._order_line_list diff --git a/examples/order-management/python-lib-core/pyproject.toml b/examples/order-management/python-lib-core/pyproject.toml index adfd0b1..d25c271 100644 --- a/examples/order-management/python-lib-core/pyproject.toml +++ b/examples/order-management/python-lib-core/pyproject.toml @@ -2,7 +2,7 @@ name = "order-management-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.7", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] 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 1379bb4..12fd607 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,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -402,35 +398,35 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_customer_list(self): from requests.customer_request import CustomerRequest @@ -483,13 +479,13 @@ def have_no_customers(self): return self.without_customer_list_matching(CustomerRequest()) def with_customer_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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 @@ -500,13 +496,13 @@ def have_no_order_statuses(self): return self.without_order_status_list_matching(OrderStatusRequest()) def with_order_status_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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 @@ -517,13 +513,13 @@ def have_no_customer_orders(self): return self.without_customer_order_list_matching(CustomerOrderRequest()) def with_customer_order_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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 @@ -534,13 +530,13 @@ def have_no_products(self): return self.without_product_list_matching(ProductRequest()) def with_product_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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 @@ -551,13 +547,13 @@ def have_no_order_lines(self): return self.without_order_line_list_matching(OrderLineRequest()) def with_order_line_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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 @@ -568,13 +564,13 @@ def have_no_order_search_presets(self): return self.without_order_search_preset_list_matching(OrderSearchPresetRequest()) def with_order_search_preset_list_matching(self, child_request): + child_request.query.projection = ["commerce_platform"] 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): + child_request.query.projection = ["commerce_platform"] 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") @@ -585,7 +581,9 @@ def count_customers_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("customer_list", alias, child_request.query, True) + ) return self @@ -598,7 +596,9 @@ def count_order_statuses_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("order_status_list", alias, child_request.query, True) + ) return self def min_display_order_of_order_statuses(self): @@ -607,8 +607,10 @@ def min_display_order_of_order_statuses(self): "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) + child_request.query.min("display_order", "min_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -616,8 +618,10 @@ def max_display_order_of_order_statuses(self): "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) + child_request.query.max("display_order", "max_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -625,8 +629,10 @@ def sum_display_order_of_order_statuses(self): "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) + child_request.query.sum("display_order", "sum_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -634,8 +640,10 @@ def avg_display_order_of_order_statuses(self): "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) + child_request.query.avg("display_order", "avg_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -643,8 +651,10 @@ def standardDeviation_display_order_of_order_statuses(self): "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) + child_request.query.standardDeviation("display_order", "standardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -652,8 +662,10 @@ def squareRootOfPopulationStandardDeviation_display_order_of_order_statuses(self "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) + child_request.query.squareRootOfPopulationStandardDeviation("display_order", "squareRootOfPopulationStandardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -661,8 +673,10 @@ def sampleVariance_display_order_of_order_statuses(self): "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) + child_request.query.sampleVariance("display_order", "sampleVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -670,8 +684,10 @@ def samplePopulationVariance_display_order_of_order_statuses(self): "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) + child_request.query.samplePopulationVariance("display_order", "samplePopulationVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("order_status_list", alias, child_request.query, True) + ) return self def count_customer_orders(self): return self.count_customer_orders_as("count_customer_orders") @@ -682,7 +698,9 @@ def count_customer_orders_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("customer_order_list", alias, child_request.query, True) + ) return self def min_total_amount_of_customer_orders(self): @@ -691,8 +709,10 @@ def min_total_amount_of_customer_orders(self): "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) + child_request.query.min("total_amount", "min_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -700,8 +720,10 @@ def max_total_amount_of_customer_orders(self): "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) + child_request.query.max("total_amount", "max_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -709,8 +731,10 @@ def sum_total_amount_of_customer_orders(self): "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) + child_request.query.sum("total_amount", "sum_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -718,8 +742,10 @@ def avg_total_amount_of_customer_orders(self): "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) + child_request.query.avg("total_amount", "avg_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -727,8 +753,10 @@ def standardDeviation_total_amount_of_customer_orders(self): "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) + child_request.query.standardDeviation("total_amount", "standardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -736,8 +764,10 @@ def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self "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) + child_request.query.squareRootOfPopulationStandardDeviation("total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -745,8 +775,10 @@ def sampleVariance_total_amount_of_customer_orders(self): "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) + child_request.query.sampleVariance("total_amount", "sampleVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -754,8 +786,10 @@ def samplePopulationVariance_total_amount_of_customer_orders(self): "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) + child_request.query.samplePopulationVariance("total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("customer_order_list", alias, child_request.query, True) + ) return self def count_products(self): return self.count_products_as("count_products") @@ -766,7 +800,9 @@ def count_products_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("product_list", alias, child_request.query, True) + ) return self @@ -779,7 +815,9 @@ def count_order_lines_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def min_quantity_of_order_lines(self): @@ -788,8 +826,10 @@ def min_quantity_of_order_lines(self): "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) + child_request.query.min("quantity", "min_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def max_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -797,8 +837,10 @@ def max_quantity_of_order_lines(self): "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) + child_request.query.max("quantity", "max_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sum_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -806,8 +848,10 @@ def sum_quantity_of_order_lines(self): "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) + child_request.query.sum("quantity", "sum_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def avg_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -815,8 +859,10 @@ def avg_quantity_of_order_lines(self): "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) + child_request.query.avg("quantity", "avg_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def standardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -824,8 +870,10 @@ def standardDeviation_quantity_of_order_lines(self): "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) + child_request.query.standardDeviation("quantity", "standardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -833,8 +881,10 @@ def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): "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) + child_request.query.squareRootOfPopulationStandardDeviation("quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sampleVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -842,8 +892,10 @@ def sampleVariance_quantity_of_order_lines(self): "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) + child_request.query.sampleVariance("quantity", "sampleVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -851,8 +903,10 @@ def samplePopulationVariance_quantity_of_order_lines(self): "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) + child_request.query.samplePopulationVariance("quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("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") @@ -863,7 +917,9 @@ def count_order_search_presets_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("order_search_preset_list", alias, child_request.query, True) + ) return self @@ -890,7 +946,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -912,21 +968,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -945,6 +1001,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 f7cc51c..8dcc698 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,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -124,12 +120,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_status_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("status"), "OrderStatus", child_request.query)) return self @@ -141,12 +137,12 @@ 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"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("customer"), "Customer", child_request.query)) return self @@ -158,12 +154,12 @@ 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"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -600,119 +596,119 @@ def min_total_amount(self): return self.min_total_amount_as("minOfTotalAmount") def min_total_amount_as(self, ret_name: str): - self.query.aggregate("min", "total_amount", ret_name) + self.query.min("total_amount", ret_name) return self def max_total_amount(self): return self.max_total_amount_as("maxOfTotalAmount") def max_total_amount_as(self, ret_name: str): - self.query.aggregate("max", "total_amount", ret_name) + self.query.max("total_amount", ret_name) return self def sum_total_amount(self): return self.sum_total_amount_as("sumOfTotalAmount") def sum_total_amount_as(self, ret_name: str): - self.query.aggregate("sum", "total_amount", ret_name) + self.query.sum("total_amount", ret_name) return self def avg_total_amount(self): return self.avg_total_amount_as("avgOfTotalAmount") def avg_total_amount_as(self, ret_name: str): - self.query.aggregate("avg", "total_amount", ret_name) + self.query.avg("total_amount", ret_name) return self def standardDeviation_total_amount(self): return self.standardDeviation_total_amount_as("standardDeviationOfTotalAmount") def standardDeviation_total_amount_as(self, ret_name: str): - self.query.aggregate("stddev", "total_amount", ret_name) + self.query.standardDeviation("total_amount", ret_name) return self def squareRootOfPopulationStandardDeviation_total_amount(self): return self.squareRootOfPopulationStandardDeviation_total_amount_as("squareRootOfPopulationStandardDeviationOfTotalAmount") def squareRootOfPopulationStandardDeviation_total_amount_as(self, ret_name: str): - self.query.aggregate("stddev_pop", "total_amount", ret_name) + self.query.squareRootOfPopulationStandardDeviation("total_amount", ret_name) return self def sampleVariance_total_amount(self): return self.sampleVariance_total_amount_as("sampleVarianceOfTotalAmount") def sampleVariance_total_amount_as(self, ret_name: str): - self.query.aggregate("var_samp", "total_amount", ret_name) + self.query.sampleVariance("total_amount", ret_name) return self def samplePopulationVariance_total_amount(self): return self.samplePopulationVariance_total_amount_as("samplePopulationVarianceOfTotalAmount") def samplePopulationVariance_total_amount_as(self, ret_name: str): - self.query.aggregate("var_pop", "total_amount", ret_name) + self.query.samplePopulationVariance("total_amount", ret_name) return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_order_number(self): self.query.group_by("order_number") return self def group_by_order_number_as(self, ret_name: str): - self.query.group_by("order_number") + self.query.group_by("order_number") return self def group_by_order_date(self): self.query.group_by("order_date") return self def group_by_order_date_as(self, ret_name: str): - self.query.group_by("order_date") + self.query.group_by("order_date") return self def group_by_total_amount(self): self.query.group_by("total_amount") return self def group_by_total_amount_as(self, ret_name: str): - self.query.group_by("total_amount") + self.query.group_by("total_amount") return self def group_by_status(self): self.query.group_by("status") return self def group_by_status_as(self, ret_name: str): - self.query.group_by("status") + self.query.group_by("status") return self def group_by_customer(self): self.query.group_by("customer") return self def group_by_customer_as(self, ret_name: str): - self.query.group_by("customer") + self.query.group_by("customer") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_order_line_list(self): from requests.order_line_request import OrderLineRequest @@ -730,13 +726,13 @@ def have_no_order_lines(self): return self.without_order_line_list_matching(OrderLineRequest()) def with_order_line_list_matching(self, child_request): + child_request.query.projection = ["customer_order"] 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): + child_request.query.projection = ["customer_order"] 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") @@ -747,7 +743,9 @@ def count_order_lines_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def min_quantity_of_order_lines(self): @@ -756,8 +754,10 @@ def min_quantity_of_order_lines(self): "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) + child_request.query.min("quantity", "min_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def max_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -765,8 +765,10 @@ def max_quantity_of_order_lines(self): "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) + child_request.query.max("quantity", "max_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sum_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -774,8 +776,10 @@ def sum_quantity_of_order_lines(self): "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) + child_request.query.sum("quantity", "sum_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def avg_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -783,8 +787,10 @@ def avg_quantity_of_order_lines(self): "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) + child_request.query.avg("quantity", "avg_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def standardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -792,8 +798,10 @@ def standardDeviation_quantity_of_order_lines(self): "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) + child_request.query.standardDeviation("quantity", "standardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -801,8 +809,10 @@ def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): "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) + child_request.query.squareRootOfPopulationStandardDeviation("quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sampleVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -810,8 +820,10 @@ def sampleVariance_quantity_of_order_lines(self): "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) + child_request.query.sampleVariance("quantity", "sampleVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -819,8 +831,10 @@ def samplePopulationVariance_quantity_of_order_lines(self): "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) + child_request.query.samplePopulationVariance("quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("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): @@ -860,7 +874,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -882,21 +896,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -915,6 +929,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 02c8e18..868c902 100644 --- a/examples/order-management/python-lib-core/requests/customer_request.py +++ b/examples/order-management/python-lib-core/requests/customer_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -110,12 +106,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_commerce_platform_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -512,49 +508,49 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_email(self): self.query.group_by("email") return self def group_by_email_as(self, ret_name: str): - self.query.group_by("email") + self.query.group_by("email") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_customer_order_list(self): from requests.customer_order_request import CustomerOrderRequest @@ -572,13 +568,13 @@ def have_no_customer_orders(self): return self.without_customer_order_list_matching(CustomerOrderRequest()) def with_customer_order_list_matching(self, child_request): + child_request.query.projection = ["customer"] 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): + child_request.query.projection = ["customer"] 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") @@ -589,7 +585,9 @@ def count_customer_orders_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("customer_order_list", alias, child_request.query, True) + ) return self def min_total_amount_of_customer_orders(self): @@ -598,8 +596,10 @@ def min_total_amount_of_customer_orders(self): "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) + child_request.query.min("total_amount", "min_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -607,8 +607,10 @@ def max_total_amount_of_customer_orders(self): "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) + child_request.query.max("total_amount", "max_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -616,8 +618,10 @@ def sum_total_amount_of_customer_orders(self): "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) + child_request.query.sum("total_amount", "sum_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -625,8 +629,10 @@ def avg_total_amount_of_customer_orders(self): "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) + child_request.query.avg("total_amount", "avg_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -634,8 +640,10 @@ def standardDeviation_total_amount_of_customer_orders(self): "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) + child_request.query.standardDeviation("total_amount", "standardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -643,8 +651,10 @@ def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self "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) + child_request.query.squareRootOfPopulationStandardDeviation("total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -652,8 +662,10 @@ def sampleVariance_total_amount_of_customer_orders(self): "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) + child_request.query.sampleVariance("total_amount", "sampleVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -661,8 +673,10 @@ def samplePopulationVariance_total_amount_of_customer_orders(self): "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) + child_request.query.samplePopulationVariance("total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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): @@ -692,7 +706,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -714,21 +728,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -747,6 +761,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 a779f5a..135cfdb 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,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -120,12 +116,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_customer_order_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("customer_order"), "CustomerOrder", child_request.query)) return self @@ -137,12 +133,12 @@ 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"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("product"), "Product", child_request.query)) return self @@ -154,12 +150,12 @@ 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"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -565,112 +561,112 @@ def min_quantity(self): return self.min_quantity_as("minOfQuantity") def min_quantity_as(self, ret_name: str): - self.query.aggregate("min", "quantity", ret_name) + self.query.min("quantity", ret_name) return self def max_quantity(self): return self.max_quantity_as("maxOfQuantity") def max_quantity_as(self, ret_name: str): - self.query.aggregate("max", "quantity", ret_name) + self.query.max("quantity", ret_name) return self def sum_quantity(self): return self.sum_quantity_as("sumOfQuantity") def sum_quantity_as(self, ret_name: str): - self.query.aggregate("sum", "quantity", ret_name) + self.query.sum("quantity", ret_name) return self def avg_quantity(self): return self.avg_quantity_as("avgOfQuantity") def avg_quantity_as(self, ret_name: str): - self.query.aggregate("avg", "quantity", ret_name) + self.query.avg("quantity", ret_name) return self def standardDeviation_quantity(self): return self.standardDeviation_quantity_as("standardDeviationOfQuantity") def standardDeviation_quantity_as(self, ret_name: str): - self.query.aggregate("stddev", "quantity", ret_name) + self.query.standardDeviation("quantity", ret_name) return self def squareRootOfPopulationStandardDeviation_quantity(self): return self.squareRootOfPopulationStandardDeviation_quantity_as("squareRootOfPopulationStandardDeviationOfQuantity") def squareRootOfPopulationStandardDeviation_quantity_as(self, ret_name: str): - self.query.aggregate("stddev_pop", "quantity", ret_name) + self.query.squareRootOfPopulationStandardDeviation("quantity", ret_name) return self def sampleVariance_quantity(self): return self.sampleVariance_quantity_as("sampleVarianceOfQuantity") def sampleVariance_quantity_as(self, ret_name: str): - self.query.aggregate("var_samp", "quantity", ret_name) + self.query.sampleVariance("quantity", ret_name) return self def samplePopulationVariance_quantity(self): return self.samplePopulationVariance_quantity_as("samplePopulationVarianceOfQuantity") def samplePopulationVariance_quantity_as(self, ret_name: str): - self.query.aggregate("var_pop", "quantity", ret_name) + self.query.samplePopulationVariance("quantity", ret_name) return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_customer_order(self): self.query.group_by("customer_order") return self def group_by_customer_order_as(self, ret_name: str): - self.query.group_by("customer_order") + self.query.group_by("customer_order") return self def group_by_product(self): self.query.group_by("product") return self def group_by_product_as(self, ret_name: str): - self.query.group_by("product") + self.query.group_by("product") return self def group_by_product_name(self): self.query.group_by("product_name") return self def group_by_product_name_as(self, ret_name: str): - self.query.group_by("product_name") + self.query.group_by("product_name") return self def group_by_sku(self): self.query.group_by("sku") return self def group_by_sku_as(self, ret_name: str): - self.query.group_by("sku") + self.query.group_by("sku") return self def group_by_quantity(self): self.query.group_by("quantity") return self def group_by_quantity_as(self, ret_name: str): - self.query.group_by("quantity") + self.query.group_by("quantity") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def facet_by_customer_order_as(self, name: str, request: QuerySelection, include_all_facets: bool = True): @@ -710,7 +706,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -732,21 +728,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -765,6 +761,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 04e0468..f41b3e6 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,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -118,12 +114,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_commerce_platform_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -678,63 +674,63 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_filter_json(self): self.query.group_by("filter_json") return self def group_by_filter_json_as(self, ret_name: str): - self.query.group_by("filter_json") + self.query.group_by("filter_json") return self def group_by_request_id(self): self.query.group_by("request_id") return self def group_by_request_id_as(self, ret_name: str): - self.query.group_by("request_id") + self.query.group_by("request_id") return self def group_by_owner_user_id(self): self.query.group_by("owner_user_id") return self def group_by_owner_user_id_as(self, ret_name: str): - self.query.group_by("owner_user_id") + self.query.group_by("owner_user_id") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def facet_by_commerce_platform_as(self, name: str, request: QuerySelection, include_all_facets: bool = True): @@ -764,7 +760,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -786,21 +782,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -819,6 +815,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 4278758..27a10d6 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,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -110,12 +106,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_commerce_platform_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -538,98 +534,98 @@ def min_display_order(self): return self.min_display_order_as("minOfDisplayOrder") def min_display_order_as(self, ret_name: str): - self.query.aggregate("min", "display_order", ret_name) + self.query.min("display_order", ret_name) return self def max_display_order(self): return self.max_display_order_as("maxOfDisplayOrder") def max_display_order_as(self, ret_name: str): - self.query.aggregate("max", "display_order", ret_name) + self.query.max("display_order", ret_name) return self def sum_display_order(self): return self.sum_display_order_as("sumOfDisplayOrder") def sum_display_order_as(self, ret_name: str): - self.query.aggregate("sum", "display_order", ret_name) + self.query.sum("display_order", ret_name) return self def avg_display_order(self): return self.avg_display_order_as("avgOfDisplayOrder") def avg_display_order_as(self, ret_name: str): - self.query.aggregate("avg", "display_order", ret_name) + self.query.avg("display_order", ret_name) return self def standardDeviation_display_order(self): return self.standardDeviation_display_order_as("standardDeviationOfDisplayOrder") def standardDeviation_display_order_as(self, ret_name: str): - self.query.aggregate("stddev", "display_order", ret_name) + self.query.standardDeviation("display_order", ret_name) return self def squareRootOfPopulationStandardDeviation_display_order(self): return self.squareRootOfPopulationStandardDeviation_display_order_as("squareRootOfPopulationStandardDeviationOfDisplayOrder") def squareRootOfPopulationStandardDeviation_display_order_as(self, ret_name: str): - self.query.aggregate("stddev_pop", "display_order", ret_name) + self.query.squareRootOfPopulationStandardDeviation("display_order", ret_name) return self def sampleVariance_display_order(self): return self.sampleVariance_display_order_as("sampleVarianceOfDisplayOrder") def sampleVariance_display_order_as(self, ret_name: str): - self.query.aggregate("var_samp", "display_order", ret_name) + self.query.sampleVariance("display_order", ret_name) return self def samplePopulationVariance_display_order(self): return self.samplePopulationVariance_display_order_as("samplePopulationVarianceOfDisplayOrder") def samplePopulationVariance_display_order_as(self, ret_name: str): - self.query.aggregate("var_pop", "display_order", ret_name) + self.query.samplePopulationVariance("display_order", ret_name) return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_code(self): self.query.group_by("code") return self def group_by_code_as(self, ret_name: str): - self.query.group_by("code") + self.query.group_by("code") return self def group_by_color(self): self.query.group_by("color") return self def group_by_color_as(self, ret_name: str): - self.query.group_by("color") + self.query.group_by("color") return self def group_by_display_order(self): self.query.group_by("display_order") return self def group_by_display_order_as(self, ret_name: str): - self.query.group_by("display_order") + self.query.group_by("display_order") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_customer_order_list(self): from requests.customer_order_request import CustomerOrderRequest @@ -647,13 +643,13 @@ def have_no_customer_orders(self): return self.without_customer_order_list_matching(CustomerOrderRequest()) def with_customer_order_list_matching(self, child_request): + child_request.query.projection = ["status"] 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): + child_request.query.projection = ["status"] 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") @@ -664,7 +660,9 @@ def count_customer_orders_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("customer_order_list", alias, child_request.query, True) + ) return self def min_total_amount_of_customer_orders(self): @@ -673,8 +671,10 @@ def min_total_amount_of_customer_orders(self): "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) + child_request.query.min("total_amount", "min_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -682,8 +682,10 @@ def max_total_amount_of_customer_orders(self): "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) + child_request.query.max("total_amount", "max_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -691,8 +693,10 @@ def sum_total_amount_of_customer_orders(self): "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) + child_request.query.sum("total_amount", "sum_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -700,8 +704,10 @@ def avg_total_amount_of_customer_orders(self): "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) + child_request.query.avg("total_amount", "avg_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -709,8 +715,10 @@ def standardDeviation_total_amount_of_customer_orders(self): "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) + child_request.query.standardDeviation("total_amount", "standardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -718,8 +726,10 @@ def squareRootOfPopulationStandardDeviation_total_amount_of_customer_orders(self "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) + child_request.query.squareRootOfPopulationStandardDeviation("total_amount", "squareRootOfPopulationStandardDeviation_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -727,8 +737,10 @@ def sampleVariance_total_amount_of_customer_orders(self): "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) + child_request.query.sampleVariance("total_amount", "sampleVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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 @@ -736,8 +748,10 @@ def samplePopulationVariance_total_amount_of_customer_orders(self): "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) + child_request.query.samplePopulationVariance("total_amount", "samplePopulationVariance_total_amount") + self.query.relation_aggregates.append( + RelationAggregate("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): @@ -767,7 +781,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -789,21 +803,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -822,6 +836,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 9ee97e0..5bbbad0 100644 --- a/examples/order-management/python-lib-core/requests/product_request.py +++ b/examples/order-management/python-lib-core/requests/product_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -65,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -114,12 +110,12 @@ def select_commerce_platform_with(self, child_request): self.query.relation_query("commerce_platform", child_request.query) return self def with_commerce_platform_matching(self, child_request): - child_request.query._projection = ["id"] + 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"] + child_request.query.projection = ["id"] self.query.and_filter(not_in_subquery(column("commerce_platform"), "CommercePlatform", child_request.query)) return self @@ -595,56 +591,56 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_sku(self): self.query.group_by("sku") return self def group_by_sku_as(self, ret_name: str): - self.query.group_by("sku") + self.query.group_by("sku") return self def group_by_image_url(self): self.query.group_by("image_url") return self def group_by_image_url_as(self, ret_name: str): - self.query.group_by("image_url") + self.query.group_by("image_url") return self def group_by_commerce_platform(self): self.query.group_by("commerce_platform") return self def group_by_commerce_platform_as(self, ret_name: str): - self.query.group_by("commerce_platform") + self.query.group_by("commerce_platform") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_order_line_list(self): from requests.order_line_request import OrderLineRequest @@ -662,13 +658,13 @@ def have_no_order_lines(self): return self.without_order_line_list_matching(OrderLineRequest()) def with_order_line_list_matching(self, child_request): + child_request.query.projection = ["product"] 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): + child_request.query.projection = ["product"] 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") @@ -679,7 +675,9 @@ def count_order_lines_as(self, alias: str): 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) + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def min_quantity_of_order_lines(self): @@ -688,8 +686,10 @@ def min_quantity_of_order_lines(self): "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) + child_request.query.min("quantity", "min_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def max_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -697,8 +697,10 @@ def max_quantity_of_order_lines(self): "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) + child_request.query.max("quantity", "max_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sum_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -706,8 +708,10 @@ def sum_quantity_of_order_lines(self): "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) + child_request.query.sum("quantity", "sum_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def avg_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -715,8 +719,10 @@ def avg_quantity_of_order_lines(self): "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) + child_request.query.avg("quantity", "avg_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def standardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -724,8 +730,10 @@ def standardDeviation_quantity_of_order_lines(self): "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) + child_request.query.standardDeviation("quantity", "standardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -733,8 +741,10 @@ def squareRootOfPopulationStandardDeviation_quantity_of_order_lines(self): "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) + child_request.query.squareRootOfPopulationStandardDeviation("quantity", "squareRootOfPopulationStandardDeviation_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def sampleVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -742,8 +752,10 @@ def sampleVariance_quantity_of_order_lines(self): "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) + child_request.query.sampleVariance("quantity", "sampleVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("order_line_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_quantity_of_order_lines(self): from requests.order_line_request import OrderLineRequest @@ -751,8 +763,10 @@ def samplePopulationVariance_quantity_of_order_lines(self): "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) + child_request.query.samplePopulationVariance("quantity", "samplePopulationVariance_quantity") + self.query.relation_aggregates.append( + RelationAggregate("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): @@ -782,7 +796,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -804,21 +818,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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) @@ -837,6 +851,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), 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 index 4979871..b53d7c2 100644 --- a/examples/order-management/python-lib-core/runtime_module.py +++ b/examples/order-management/python-lib-core/runtime_module.py @@ -293,31 +293,38 @@ def check_and_fix(self, context, record, location, results): _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()) + .audit_mask_fields([]) + .table_name("commerce_platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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()) + .audit_mask_fields(["email"]) + .table_name("customer_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("email", DataType.Text).column_name("email").log_policy("plain").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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()) + .audit_mask_fields([]) + .table_name("order_status_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("code", DataType.Text).column_name("code").log_policy("plain").required()).property(PropertyDescriptor("color", DataType.Text).column_name("color").log_policy("plain")).property(PropertyDescriptor("display_order", DataType.Decimal).column_name("display_order").log_policy("plain")).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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()) + .audit_mask_fields([]) + .table_name("customer_order_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("order_number", DataType.Text).column_name("order_number").log_policy("plain").required()).property(PropertyDescriptor("order_date", DataType.Date).column_name("order_date").log_policy("plain").required()).property(PropertyDescriptor("total_amount", DataType.Decimal).column_name("total_amount").log_policy("plain").required()).property(PropertyDescriptor("status", DataType.I64).column_name("status").log_policy("plain").required()).property(PropertyDescriptor("customer", DataType.I64).column_name("customer").log_policy("plain").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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()) + .audit_mask_fields([]) + .table_name("product_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("sku", DataType.Text).column_name("sku").log_policy("plain").required()).property(PropertyDescriptor("image_url", DataType.Text).column_name("image_url").log_policy("plain")).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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")) + .audit_mask_fields([]) + .table_name("order_line_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("customer_order", DataType.I64).column_name("customer_order").log_policy("plain").required()).property(PropertyDescriptor("product", DataType.I64).column_name("product").log_policy("plain").required()).property(PropertyDescriptor("product_name", DataType.Text).column_name("product_name").log_policy("plain").required()).property(PropertyDescriptor("sku", DataType.Text).column_name("sku").log_policy("plain").required()).property(PropertyDescriptor("quantity", DataType.I64).column_name("quantity").log_policy("plain").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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")) + .audit_mask_fields([]) + .table_name("order_search_preset_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("filter_json", DataType.Text).column_name("filter_json").log_policy("plain").required()).property(PropertyDescriptor("request_id", DataType.Text).column_name("request_id").log_policy("plain").required()).property(PropertyDescriptor("owner_user_id", DataType.Text).column_name("owner_user_id").log_policy("plain").required()).property(PropertyDescriptor("commerce_platform", DataType.I64).column_name("commerce_platform").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("commerce_platform", "CommercePlatform").local("commerce_platform").foreign("id")) ) async def _ensure_generated_bootstrap_once(context): diff --git a/examples/order-management/test_sql_log_intent.py b/examples/order-management/test_sql_log_intent.py new file mode 100644 index 0000000..b5f9ed9 --- /dev/null +++ b/examples/order-management/test_sql_log_intent.py @@ -0,0 +1,37 @@ +"""Check the regenerated order library's SQL intent through its real app.""" +import os +from pathlib import Path +import subprocess +import sys +import tempfile +import unittest + + +class OrderSqlLogIntentTest(unittest.TestCase): + def test_generated_queries_retain_intent_at_sql_sink(self): + repo = Path(__file__).resolve().parents[2] + env = dict(os.environ) + env.pop("TEAQL_ALLOW_SENSITIVE_PLAINTEXT_LOGS", None) + env["PYTHONPATH"] = os.pathsep.join((str(repo / "examples" / "order-management" / "python-lib-core"), str(repo / "src"))) + with tempfile.TemporaryDirectory(prefix="teaql-order-log-test-") as directory: + env["TEAQL_ORDER_MANAGEMENT_DB"] = str(Path(directory) / "order.sqlite") + result = subprocess.run([sys.executable, str(Path(__file__).with_name("python-app-console") / "app.py")], + cwd=repo, env=env, capture_output=True, text=True, timeout=90) + output = result.stdout + result.stderr + self.assertEqual(0, result.returncode, output) + queries = [line for line in output.splitlines() + if line.startswith("[TeaQL SQL]") and "[select]" in line] + self.assertGreater(len(queries), 0, output) + for line in queries: + self.assertNotIn("comment=None", line) + self.assertNotIn("purpose=None", line) + self.assertNotIn("masked-in-quick-start", output) + customer_sql = next((line for line in output.splitlines() + if "INSERT INTO customer_data" in line), "") + self.assertIn("/* masked */", customer_sql) + self.assertIn("'Acme Retail'", customer_sql) + self.assertIn("[schema] ensured 7 generated entity tables", output) + + +if __name__ == "__main__": + unittest.main() diff --git a/examples/school-management/app/main.py b/examples/school-management/app/main.py index 9225aad..33bd0bb 100644 --- a/examples/school-management/app/main.py +++ b/examples/school-management/app/main.py @@ -1,5 +1,6 @@ import asyncio from datetime import date +import os from pathlib import Path import sys @@ -50,7 +51,7 @@ async def verify_dynamic_search(context): async def main() -> None: - database = ROOT / ".local" / "school.sqlite" + database = Path(os.environ.get("TEAQL_SCHOOL_MANAGEMENT_DB", ROOT / ".local" / "school.sqlite")) database.parent.mkdir(parents=True, exist_ok=True) database.unlink(missing_ok=True) client = SQLiteTeaQLClient(str(database)) diff --git a/examples/school-management/models/platform.py b/examples/school-management/models/platform.py index 3e2eb58..a2878ec 100644 --- a/examples/school-management/models/platform.py +++ b/examples/school-management/models/platform.py @@ -333,4 +333,4 @@ def school_type_list(self) -> list: def school_list(self) -> list: self._loaded_fields.add("school_list") - return self._school_list \ No newline at end of file + return self._school_list diff --git a/examples/school-management/models/school.py b/examples/school-management/models/school.py index 0634f6e..43d2345 100644 --- a/examples/school-management/models/school.py +++ b/examples/school-management/models/school.py @@ -379,5 +379,5 @@ def update_school_type(self, value): def update_school_type_to_primary(self): self.schoolType = 1001 self._loaded_fields.add("schoolType") + self._entity_root.set(self._teaql_entity_key(), "school_type", Value.from_any(self.schoolType)) return self - diff --git a/examples/school-management/models/school_type.py b/examples/school-management/models/school_type.py index ed64642..418f288 100644 --- a/examples/school-management/models/school_type.py +++ b/examples/school-management/models/school_type.py @@ -306,4 +306,4 @@ def update_platform(self, value): def school_list(self) -> list: self._loaded_fields.add("school_list") - return self._school_list \ No newline at end of file + return self._school_list diff --git a/examples/school-management/pyproject.toml b/examples/school-management/pyproject.toml index a7537f1..7984d34 100644 --- a/examples/school-management/pyproject.toml +++ b/examples/school-management/pyproject.toml @@ -2,7 +2,7 @@ name = "school-management-service-lib" version = "1.0.0" description = "Generated python library" -dependencies = ["teaql==0.2.5", "aiosqlite>=0.22.1"] +dependencies = ["teaql==0.2.7", "aiosqlite>=0.22.1"] [tool.setuptools] py-modules = ["Q", "E"] diff --git a/examples/school-management/requests/platform_request.py b/examples/school-management/requests/platform_request.py index 3f59c85..503e9ec 100644 --- a/examples/school-management/requests/platform_request.py +++ b/examples/school-management/requests/platform_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -10,7 +10,6 @@ ) from models.platform import Platform from typing import Protocol -from copy import deepcopy class QuerySelection(Protocol): query: SelectQuery @@ -66,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -486,42 +481,42 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_base_url(self): self.query.group_by("base_url") return self def group_by_base_url_as(self, ret_name: str): - self.query.group_by("base_url") + self.query.group_by("base_url") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_school_type_list(self): from requests.school_type_request import SchoolTypeRequest @@ -546,15 +541,13 @@ def have_no_school_types(self): return self.without_school_type_list_matching(SchoolTypeRequest()) def with_school_type_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["platform"] - self.query.and_filter(in_subquery(column("id"), "SchoolType", child_query)) + child_request.query.projection = ["platform"] + self.query.and_filter(in_subquery(column("id"), "SchoolType", child_request.query)) return self def without_school_type_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["platform"] - self.query.and_filter(not_in_subquery(column("id"), "SchoolType", child_query)) + child_request.query.projection = ["platform"] + self.query.and_filter(not_in_subquery(column("id"), "SchoolType", child_request.query)) return self def have_schools(self): from requests.school_request import SchoolRequest @@ -565,15 +558,13 @@ def have_no_schools(self): return self.without_school_list_matching(SchoolRequest()) def with_school_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["platform"] - self.query.and_filter(in_subquery(column("id"), "School", child_query)) + child_request.query.projection = ["platform"] + self.query.and_filter(in_subquery(column("id"), "School", child_request.query)) return self def without_school_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["platform"] - self.query.and_filter(not_in_subquery(column("id"), "School", child_query)) + child_request.query.projection = ["platform"] + self.query.and_filter(not_in_subquery(column("id"), "School", child_request.query)) return self def count_school_types(self): return self.count_school_types_as("count_school_types") @@ -584,7 +575,9 @@ def count_school_types_as(self, alias: str): def count_school_types_with(self, alias: str, child_request): child_request.query.count_field("id", alias) - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def min_display_order_of_school_types(self): @@ -593,8 +586,10 @@ def min_display_order_of_school_types(self): "min_display_order_of_school_types", SchoolTypeRequest()) def min_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("min", "display_order", "min_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.min("display_order", "min_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def max_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -602,8 +597,10 @@ def max_display_order_of_school_types(self): "max_display_order_of_school_types", SchoolTypeRequest()) def max_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("max", "display_order", "max_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.max("display_order", "max_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def sum_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -611,8 +608,10 @@ def sum_display_order_of_school_types(self): "sum_display_order_of_school_types", SchoolTypeRequest()) def sum_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("sum", "display_order", "sum_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.sum("display_order", "sum_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def avg_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -620,8 +619,10 @@ def avg_display_order_of_school_types(self): "avg_display_order_of_school_types", SchoolTypeRequest()) def avg_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("avg", "display_order", "avg_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.avg("display_order", "avg_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def standardDeviation_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -629,8 +630,10 @@ def standardDeviation_display_order_of_school_types(self): "standardDeviation_display_order_of_school_types", SchoolTypeRequest()) def standardDeviation_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("stddev", "display_order", "standardDeviation_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.standardDeviation("display_order", "standardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -638,8 +641,10 @@ def squareRootOfPopulationStandardDeviation_display_order_of_school_types(self): "squareRootOfPopulationStandardDeviation_display_order_of_school_types", SchoolTypeRequest()) def squareRootOfPopulationStandardDeviation_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("stddev_pop", "display_order", "squareRootOfPopulationStandardDeviation_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.squareRootOfPopulationStandardDeviation("display_order", "squareRootOfPopulationStandardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def sampleVariance_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -647,8 +652,10 @@ def sampleVariance_display_order_of_school_types(self): "sampleVariance_display_order_of_school_types", SchoolTypeRequest()) def sampleVariance_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("var_samp", "display_order", "sampleVariance_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.sampleVariance("display_order", "sampleVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_display_order_of_school_types(self): from requests.school_type_request import SchoolTypeRequest @@ -656,8 +663,10 @@ def samplePopulationVariance_display_order_of_school_types(self): "samplePopulationVariance_display_order_of_school_types", SchoolTypeRequest()) def samplePopulationVariance_display_order_of_school_types_as(self, alias: str, child_request): - child_request.query.aggregate("var_pop", "display_order", "samplePopulationVariance_display_order") - self.query.relation_aggregate("school_type_list", alias, child_request.query, True) + child_request.query.samplePopulationVariance("display_order", "samplePopulationVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("school_type_list", alias, child_request.query, True) + ) return self def count_schools(self): return self.count_schools_as("count_schools") @@ -668,7 +677,9 @@ def count_schools_as(self, alias: str): def count_schools_with(self, alias: str, child_request): child_request.query.count_field("id", alias) - self.query.relation_aggregate("school_list", alias, child_request.query, True) + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def min_student_capacity_of_schools(self): @@ -677,8 +688,10 @@ def min_student_capacity_of_schools(self): "min_student_capacity_of_schools", SchoolRequest()) def min_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("min", "student_capacity", "min_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.min("student_capacity", "min_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def max_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -686,8 +699,10 @@ def max_student_capacity_of_schools(self): "max_student_capacity_of_schools", SchoolRequest()) def max_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("max", "student_capacity", "max_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.max("student_capacity", "max_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def sum_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -695,8 +710,10 @@ def sum_student_capacity_of_schools(self): "sum_student_capacity_of_schools", SchoolRequest()) def sum_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("sum", "student_capacity", "sum_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.sum("student_capacity", "sum_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def avg_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -704,8 +721,10 @@ def avg_student_capacity_of_schools(self): "avg_student_capacity_of_schools", SchoolRequest()) def avg_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("avg", "student_capacity", "avg_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.avg("student_capacity", "avg_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def standardDeviation_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -713,8 +732,10 @@ def standardDeviation_student_capacity_of_schools(self): "standardDeviation_student_capacity_of_schools", SchoolRequest()) def standardDeviation_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("stddev", "student_capacity", "standardDeviation_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.standardDeviation("student_capacity", "standardDeviation_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -722,8 +743,10 @@ def squareRootOfPopulationStandardDeviation_student_capacity_of_schools(self): "squareRootOfPopulationStandardDeviation_student_capacity_of_schools", SchoolRequest()) def squareRootOfPopulationStandardDeviation_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("stddev_pop", "student_capacity", "squareRootOfPopulationStandardDeviation_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.squareRootOfPopulationStandardDeviation("student_capacity", "squareRootOfPopulationStandardDeviation_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def sampleVariance_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -731,8 +754,10 @@ def sampleVariance_student_capacity_of_schools(self): "sampleVariance_student_capacity_of_schools", SchoolRequest()) def sampleVariance_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("var_samp", "student_capacity", "sampleVariance_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.sampleVariance("student_capacity", "sampleVariance_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -740,8 +765,10 @@ def samplePopulationVariance_student_capacity_of_schools(self): "samplePopulationVariance_student_capacity_of_schools", SchoolRequest()) def samplePopulationVariance_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("var_pop", "student_capacity", "samplePopulationVariance_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.samplePopulationVariance("student_capacity", "samplePopulationVariance_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self class ExecutablePlatformRequest: @@ -766,7 +793,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -788,21 +815,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) query_root = EntityRoot() data = SmartList(Platform(_entity_root=query_root, **row) for row in row_result.rows) return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) @@ -821,6 +848,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): for row in chunk.rows: yield Platform(_entity_root=query_root, **row) diff --git a/examples/school-management/requests/school_request.py b/examples/school-management/requests/school_request.py index ddd0216..9d33284 100644 --- a/examples/school-management/requests/school_request.py +++ b/examples/school-management/requests/school_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -10,7 +10,6 @@ ) from models.school import School from typing import Protocol -from copy import deepcopy class QuerySelection(Protocol): query: SelectQuery @@ -66,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -128,15 +123,13 @@ def select_school_type_with(self, child_request): self.query.relation_query("school_type", child_request.query) return self def with_platform_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(in_subquery(column("platform"), "Platform", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("platform"), "Platform", child_request.query)) return self def without_platform_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_request.query)) return self def have_platform(self): @@ -147,15 +140,13 @@ def have_no_platform(self): self.query.and_filter(is_null(column("platform"))) return self def with_school_type_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(in_subquery(column("school_type"), "SchoolType", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("school_type"), "SchoolType", child_request.query)) return self def without_school_type_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(not_in_subquery(column("school_type"), "SchoolType", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("school_type"), "SchoolType", child_request.query)) return self def have_school_type(self): @@ -717,126 +708,126 @@ def min_student_capacity(self): return self.min_student_capacity_as("minOfStudentCapacity") def min_student_capacity_as(self, ret_name: str): - self.query.aggregate("min", "student_capacity", ret_name) + self.query.min("student_capacity", ret_name) return self def max_student_capacity(self): return self.max_student_capacity_as("maxOfStudentCapacity") def max_student_capacity_as(self, ret_name: str): - self.query.aggregate("max", "student_capacity", ret_name) + self.query.max("student_capacity", ret_name) return self def sum_student_capacity(self): return self.sum_student_capacity_as("sumOfStudentCapacity") def sum_student_capacity_as(self, ret_name: str): - self.query.aggregate("sum", "student_capacity", ret_name) + self.query.sum("student_capacity", ret_name) return self def avg_student_capacity(self): return self.avg_student_capacity_as("avgOfStudentCapacity") def avg_student_capacity_as(self, ret_name: str): - self.query.aggregate("avg", "student_capacity", ret_name) + self.query.avg("student_capacity", ret_name) return self def standardDeviation_student_capacity(self): return self.standardDeviation_student_capacity_as("standardDeviationOfStudentCapacity") def standardDeviation_student_capacity_as(self, ret_name: str): - self.query.aggregate("stddev", "student_capacity", ret_name) + self.query.standardDeviation("student_capacity", ret_name) return self def squareRootOfPopulationStandardDeviation_student_capacity(self): return self.squareRootOfPopulationStandardDeviation_student_capacity_as("squareRootOfPopulationStandardDeviationOfStudentCapacity") def squareRootOfPopulationStandardDeviation_student_capacity_as(self, ret_name: str): - self.query.aggregate("stddev_pop", "student_capacity", ret_name) + self.query.squareRootOfPopulationStandardDeviation("student_capacity", ret_name) return self def sampleVariance_student_capacity(self): return self.sampleVariance_student_capacity_as("sampleVarianceOfStudentCapacity") def sampleVariance_student_capacity_as(self, ret_name: str): - self.query.aggregate("var_samp", "student_capacity", ret_name) + self.query.sampleVariance("student_capacity", ret_name) return self def samplePopulationVariance_student_capacity(self): return self.samplePopulationVariance_student_capacity_as("samplePopulationVarianceOfStudentCapacity") def samplePopulationVariance_student_capacity_as(self, ret_name: str): - self.query.aggregate("var_pop", "student_capacity", ret_name) + self.query.samplePopulationVariance("student_capacity", ret_name) return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_platform(self): self.query.group_by("platform") return self def group_by_platform_as(self, ret_name: str): - self.query.group_by("platform") + self.query.group_by("platform") return self def group_by_school_type(self): self.query.group_by("school_type") return self def group_by_school_type_as(self, ret_name: str): - self.query.group_by("school_type") + self.query.group_by("school_type") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_address(self): self.query.group_by("address") return self def group_by_address_as(self, ret_name: str): - self.query.group_by("address") + self.query.group_by("address") return self def group_by_established_date(self): self.query.group_by("established_date") return self def group_by_established_date_as(self, ret_name: str): - self.query.group_by("established_date") + self.query.group_by("established_date") return self def group_by_student_capacity(self): self.query.group_by("student_capacity") return self def group_by_student_capacity_as(self, ret_name: str): - self.query.group_by("student_capacity") + self.query.group_by("student_capacity") return self def group_by_active(self): self.query.group_by("active") return self def group_by_active_as(self, ret_name: str): - self.query.group_by("active") + self.query.group_by("active") return self def group_by_create_time(self): self.query.group_by("create_time") return self def group_by_create_time_as(self, ret_name: str): - self.query.group_by("create_time") + self.query.group_by("create_time") return self def group_by_update_time(self): self.query.group_by("update_time") return self def group_by_update_time_as(self, ret_name: str): - self.query.group_by("update_time") + self.query.group_by("update_time") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def facet_by_platform_as(self, name: str, request: QuerySelection, include_all_facets: bool = True): @@ -871,7 +862,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -893,21 +884,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) query_root = EntityRoot() data = SmartList(School(_entity_root=query_root, **row) for row in row_result.rows) return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) @@ -926,6 +917,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): for row in chunk.rows: yield School(_entity_root=query_root, **row) diff --git a/examples/school-management/requests/school_type_request.py b/examples/school-management/requests/school_type_request.py index 7c816b1..76a097a 100644 --- a/examples/school-management/requests/school_type_request.py +++ b/examples/school-management/requests/school_type_request.py @@ -1,4 +1,4 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, SelectQuery from teaql.core.list import SmartList, TeaQLPage from teaql.runtime import EntityRoot from teaql.data_service import QueryRequest @@ -10,7 +10,6 @@ ) from models.school_type import SchoolType from typing import Protocol -from copy import deepcopy class QuerySelection(Protocol): query: SelectQuery @@ -66,15 +65,11 @@ def offset(self, n: int): return self def with_deleted_rows(self): - self.query._filters = [ - expression for expression in self.query._filters - if expression.get("field") != "version" - ] + self.query.with_deleted_rows() return self def deleted_rows_only(self): - self.with_deleted_rows() - self.query.and_filter(lte("version", -1)) + self.query.deleted_rows_only() return self def select_self_fields(self): @@ -106,15 +101,13 @@ def select_platform_with(self, child_request): self.query.relation_query("platform", child_request.query) return self def with_platform_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(in_subquery(column("platform"), "Platform", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("platform"), "Platform", child_request.query)) return self def without_platform_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["id"] - self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_query)) + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_request.query)) return self def have_platform(self): @@ -456,91 +449,91 @@ def min_display_order(self): return self.min_display_order_as("minOfDisplayOrder") def min_display_order_as(self, ret_name: str): - self.query.aggregate("min", "display_order", ret_name) + self.query.min("display_order", ret_name) return self def max_display_order(self): return self.max_display_order_as("maxOfDisplayOrder") def max_display_order_as(self, ret_name: str): - self.query.aggregate("max", "display_order", ret_name) + self.query.max("display_order", ret_name) return self def sum_display_order(self): return self.sum_display_order_as("sumOfDisplayOrder") def sum_display_order_as(self, ret_name: str): - self.query.aggregate("sum", "display_order", ret_name) + self.query.sum("display_order", ret_name) return self def avg_display_order(self): return self.avg_display_order_as("avgOfDisplayOrder") def avg_display_order_as(self, ret_name: str): - self.query.aggregate("avg", "display_order", ret_name) + self.query.avg("display_order", ret_name) return self def standardDeviation_display_order(self): return self.standardDeviation_display_order_as("standardDeviationOfDisplayOrder") def standardDeviation_display_order_as(self, ret_name: str): - self.query.aggregate("stddev", "display_order", ret_name) + self.query.standardDeviation("display_order", ret_name) return self def squareRootOfPopulationStandardDeviation_display_order(self): return self.squareRootOfPopulationStandardDeviation_display_order_as("squareRootOfPopulationStandardDeviationOfDisplayOrder") def squareRootOfPopulationStandardDeviation_display_order_as(self, ret_name: str): - self.query.aggregate("stddev_pop", "display_order", ret_name) + self.query.squareRootOfPopulationStandardDeviation("display_order", ret_name) return self def sampleVariance_display_order(self): return self.sampleVariance_display_order_as("sampleVarianceOfDisplayOrder") def sampleVariance_display_order_as(self, ret_name: str): - self.query.aggregate("var_samp", "display_order", ret_name) + self.query.sampleVariance("display_order", ret_name) return self def samplePopulationVariance_display_order(self): return self.samplePopulationVariance_display_order_as("samplePopulationVarianceOfDisplayOrder") def samplePopulationVariance_display_order_as(self, ret_name: str): - self.query.aggregate("var_pop", "display_order", ret_name) + self.query.samplePopulationVariance("display_order", ret_name) return self def group_by_platform(self): self.query.group_by("platform") return self def group_by_platform_as(self, ret_name: str): - self.query.group_by("platform") + self.query.group_by("platform") return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self def group_by_code(self): self.query.group_by("code") return self def group_by_code_as(self, ret_name: str): - self.query.group_by("code") + self.query.group_by("code") return self def group_by_display_order(self): self.query.group_by("display_order") return self def group_by_display_order_as(self, ret_name: str): - self.query.group_by("display_order") + self.query.group_by("display_order") return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") return self def select_school_list(self): from requests.school_request import SchoolRequest @@ -558,15 +551,13 @@ def have_no_schools(self): return self.without_school_list_matching(SchoolRequest()) def with_school_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["school_type"] - self.query.and_filter(in_subquery(column("id"), "School", child_query)) + child_request.query.projection = ["school_type"] + self.query.and_filter(in_subquery(column("id"), "School", child_request.query)) return self def without_school_list_matching(self, child_request): - child_query = deepcopy(child_request.query) - child_query.projection = ["school_type"] - self.query.and_filter(not_in_subquery(column("id"), "School", child_query)) + child_request.query.projection = ["school_type"] + self.query.and_filter(not_in_subquery(column("id"), "School", child_request.query)) return self def count_schools(self): return self.count_schools_as("count_schools") @@ -577,7 +568,9 @@ def count_schools_as(self, alias: str): def count_schools_with(self, alias: str, child_request): child_request.query.count_field("id", alias) - self.query.relation_aggregate("school_list", alias, child_request.query, True) + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def min_student_capacity_of_schools(self): @@ -586,8 +579,10 @@ def min_student_capacity_of_schools(self): "min_student_capacity_of_schools", SchoolRequest()) def min_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("min", "student_capacity", "min_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.min("student_capacity", "min_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def max_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -595,8 +590,10 @@ def max_student_capacity_of_schools(self): "max_student_capacity_of_schools", SchoolRequest()) def max_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("max", "student_capacity", "max_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.max("student_capacity", "max_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def sum_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -604,8 +601,10 @@ def sum_student_capacity_of_schools(self): "sum_student_capacity_of_schools", SchoolRequest()) def sum_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("sum", "student_capacity", "sum_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.sum("student_capacity", "sum_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def avg_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -613,8 +612,10 @@ def avg_student_capacity_of_schools(self): "avg_student_capacity_of_schools", SchoolRequest()) def avg_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("avg", "student_capacity", "avg_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.avg("student_capacity", "avg_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def standardDeviation_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -622,8 +623,10 @@ def standardDeviation_student_capacity_of_schools(self): "standardDeviation_student_capacity_of_schools", SchoolRequest()) def standardDeviation_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("stddev", "student_capacity", "standardDeviation_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.standardDeviation("student_capacity", "standardDeviation_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def squareRootOfPopulationStandardDeviation_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -631,8 +634,10 @@ def squareRootOfPopulationStandardDeviation_student_capacity_of_schools(self): "squareRootOfPopulationStandardDeviation_student_capacity_of_schools", SchoolRequest()) def squareRootOfPopulationStandardDeviation_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("stddev_pop", "student_capacity", "squareRootOfPopulationStandardDeviation_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.squareRootOfPopulationStandardDeviation("student_capacity", "squareRootOfPopulationStandardDeviation_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def sampleVariance_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -640,8 +645,10 @@ def sampleVariance_student_capacity_of_schools(self): "sampleVariance_student_capacity_of_schools", SchoolRequest()) def sampleVariance_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("var_samp", "student_capacity", "sampleVariance_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.sampleVariance("student_capacity", "sampleVariance_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def samplePopulationVariance_student_capacity_of_schools(self): from requests.school_request import SchoolRequest @@ -649,8 +656,10 @@ def samplePopulationVariance_student_capacity_of_schools(self): "samplePopulationVariance_student_capacity_of_schools", SchoolRequest()) def samplePopulationVariance_student_capacity_of_schools_as(self, alias: str, child_request): - child_request.query.aggregate("var_pop", "student_capacity", "samplePopulationVariance_student_capacity") - self.query.relation_aggregate("school_list", alias, child_request.query, True) + child_request.query.samplePopulationVariance("student_capacity", "samplePopulationVariance_student_capacity") + self.query.relation_aggregates.append( + RelationAggregate("school_list", alias, child_request.query, True) + ) return self def facet_by_platform_as(self, name: str, request: QuerySelection, include_all_facets: bool = True): @@ -680,7 +689,7 @@ async def execute_for_result(self, context): 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(context.prepare_query(self.query)) + req = QueryRequest(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) return await service.query(context, req) async def execute_for_rows(self, context): @@ -702,21 +711,21 @@ async def execute_for_page(self, context, offset: int, limit: int) -> TeaQLPage[ service = context.require_resource("dataService") alias = "__teaql_total" if authorized.id_set_pagination is not None: - row_result = await service.query(context, QueryRequest(authorized)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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))) + count_result = await service.query(context, QueryRequest(authorized.for_exact_count(alias), _comment=request._comment, _purpose=request._purpose)) 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)) + row_result = await service.query(context, QueryRequest(authorized, _comment=request._comment, _purpose=request._purpose)) query_root = EntityRoot() data = SmartList(SchoolType(_entity_root=query_root, **row) for row in row_result.rows) return TeaQLPage(data=data, total_count=total_count, offset=offset, limit=limit) @@ -735,6 +744,6 @@ async def execute_for_stream(self, context, chunk_size: int = 1000): 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): + async for chunk in service.query_stream(context, QueryRequest(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): for row in chunk.rows: yield SchoolType(_entity_root=query_root, **row) diff --git a/examples/school-management/runtime_module.py b/examples/school-management/runtime_module.py index cb468ae..c575ede 100644 --- a/examples/school-management/runtime_module.py +++ b/examples/school-management/runtime_module.py @@ -139,15 +139,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()) + .audit_mask_fields([]) + .table_name("platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("base_url", DataType.Text).column_name("base_url").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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()) + .audit_mask_fields([]) + .table_name("school_type_data").property(PropertyDescriptor("platform", DataType.I64).column_name("platform").log_policy("plain").required()).property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("code", DataType.Text).column_name("code").log_policy("plain").required()).property(PropertyDescriptor("display_order", DataType.Decimal).column_name("display_order").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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")) + .audit_mask_fields([]) + .table_name("school_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").log_policy("plain").required()).property(PropertyDescriptor("school_type", DataType.I64).column_name("school_type").log_policy("plain").required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("address", DataType.Text).column_name("address").log_policy("plain").required()).property(PropertyDescriptor("established_date", DataType.Date).column_name("established_date").log_policy("plain").required()).property(PropertyDescriptor("student_capacity", DataType.I64).column_name("student_capacity").log_policy("plain").required()).property(PropertyDescriptor("active", DataType.Bool).column_name("active").log_policy("plain").required()).property(PropertyDescriptor("create_time", DataType.Timestamp).column_name("create_time").log_policy("plain").required()).property(PropertyDescriptor("update_time", DataType.Timestamp).column_name("update_time").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").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): diff --git a/examples/school-management/test_sql_log_intent.py b/examples/school-management/test_sql_log_intent.py new file mode 100644 index 0000000..2db1a09 --- /dev/null +++ b/examples/school-management/test_sql_log_intent.py @@ -0,0 +1,32 @@ +"""Exercise the regenerated School library through its real application flow.""" +import os +from pathlib import Path +import subprocess +import sys +import tempfile +import unittest + + +class SchoolSqlLogIntentTest(unittest.TestCase): + def test_generated_queries_retain_intent_at_sql_sink(self): + repo = Path(__file__).resolve().parents[2] + env = dict(os.environ) + env.pop("TEAQL_ALLOW_SENSITIVE_PLAINTEXT_LOGS", None) + env["PYTHONPATH"] = os.pathsep.join((str(repo / "examples" / "school-management"), str(repo / "src"))) + with tempfile.TemporaryDirectory(prefix="teaql-school-log-test-") as directory: + env["TEAQL_SCHOOL_MANAGEMENT_DB"] = str(Path(directory) / "school.sqlite") + result = subprocess.run([sys.executable, "-m", "app.main"], cwd=repo, + env=env, capture_output=True, text=True, timeout=90) + output = result.stdout + result.stderr + self.assertEqual(0, result.returncode, output) + queries = [line for line in output.splitlines() + if line.startswith("[TeaQL SQL]") and "[select]" in line] + self.assertGreater(len(queries), 0, output) + for line in queries: + self.assertNotIn("comment=None", line) + self.assertNotIn("purpose=None", line) + self.assertIn("PASS Python School Management", output) + + +if __name__ == "__main__": + unittest.main() diff --git a/examples/task_board/generated/E.py b/examples/task_board/generated/E.py new file mode 100644 index 0000000..a801e03 --- /dev/null +++ b/examples/task_board/generated/E.py @@ -0,0 +1,217 @@ +class TeaQLNotLoadedError(RuntimeError): + def __init__(self, root, access_path, break_point): + self.root = root + self.access_path = access_path + self.break_point = break_point + super().__init__( + f"TeaQLNotLoadedError: root={root} access_path={access_path} " + f"break_point={break_point} suggested_fix=select_{break_point}(...)" + ) + + +class ValueExpression: + def __init__(self, value=None, error=None): + self._value = value + self._error = error + + def eval(self): + if self._error is not None: + raise self._error + return self._value + + def or_if_null(self, fallback): + value = self.eval() + return fallback if value is None else value + + +class EntityExpression: + def __init__(self, value, root=None, path="", error=None): + self._value = value + self._root = root or f"{type(value).__name__ if value is not None else 'Entity'}(null)" + self._path = path + self._error = error + + def eval(self): + if self._error is not None: + raise self._error + return self._value + + 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, 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(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 self._item_expression(value, self._root, path) + + +class PlatformExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + def founded(self): + return self._scalar("founded") + def user_email(self): + return self._scalar("userEmail") + def version(self): + return self._scalar("version") + def task_status_list(self): + path = self._path_for("task_status_list") + if self._error is not None: + return ListExpression([], self._root, path, TaskStatusExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, TaskStatusExpression) + if "task_status_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, TaskStatusExpression, self._not_loaded("task_status_list")) + return ListExpression(getattr(self._value, "_task_status_list"), self._root, path, TaskStatusExpression) + def task_list(self): + path = self._path_for("task_list") + if self._error is not None: + return ListExpression([], self._root, path, TaskExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, TaskExpression) + if "task_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, TaskExpression, self._not_loaded("task_list")) + return ListExpression(getattr(self._value, "_task_list"), self._root, path, TaskExpression) + pass + +class TaskStatusExpression(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 progress(self): + return self._scalar("progress") + def version(self): + return self._scalar("version") + def platform_id(self): + return self._scalar("platform", relation_id=True) + + def platform(self): + return self._relation("platform", PlatformExpression) + def task_list(self): + path = self._path_for("task_list") + if self._error is not None: + return ListExpression([], self._root, path, TaskExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, TaskExpression) + if "task_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, TaskExpression, self._not_loaded("task_list")) + return ListExpression(getattr(self._value, "_task_list"), self._root, path, TaskExpression) + pass + +class TaskExpression(EntityExpression): + def id(self): + return self._scalar("id") + def name(self): + return self._scalar("name") + 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", TaskStatusExpression) + def platform_id(self): + return self._scalar("platform", relation_id=True) + + def platform(self): + return self._relation("platform", PlatformExpression) + def task_execution_log_list(self): + path = self._path_for("task_execution_log_list") + if self._error is not None: + return ListExpression([], self._root, path, TaskExecutionLogExpression, self._error) + if self._value is None: + return ListExpression([], self._root, path, TaskExecutionLogExpression) + if "task_execution_log_list" not in getattr(self._value, "_loaded_fields", set()): + return ListExpression([], self._root, path, TaskExecutionLogExpression, self._not_loaded("task_execution_log_list")) + return ListExpression(getattr(self._value, "_task_execution_log_list"), self._root, path, TaskExecutionLogExpression) + pass + +class TaskExecutionLogExpression(EntityExpression): + def id(self): + return self._scalar("id") + def action(self): + return self._scalar("action") + def detail(self): + return self._scalar("detail") + def version(self): + return self._scalar("version") + def task_id(self): + return self._scalar("task", relation_id=True) + + def task(self): + return self._relation("task", TaskExpression) + pass + +class E: + @staticmethod + def platform(value): + entity_id = getattr(value, "id", None) + return PlatformExpression(value, "Platform(id={})".format(entity_id)) + @staticmethod + def task_status(value): + entity_id = getattr(value, "id", None) + return TaskStatusExpression(value, "TaskStatus(id={})".format(entity_id)) + @staticmethod + def task(value): + entity_id = getattr(value, "id", None) + return TaskExpression(value, "Task(id={})".format(entity_id)) + @staticmethod + def task_execution_log(value): + entity_id = getattr(value, "id", None) + return TaskExecutionLogExpression(value, "TaskExecutionLog(id={})".format(entity_id)) + pass \ No newline at end of file diff --git a/examples/task_board/generated/Q.py b/examples/task_board/generated/Q.py new file mode 100644 index 0000000..7cb255d --- /dev/null +++ b/examples/task_board/generated/Q.py @@ -0,0 +1,38 @@ +# Generated by teaql-code-gen +from requests.platform_request import PlatformRequest +from requests.task_status_request import TaskStatusRequest +from requests.task_request import TaskRequest +from requests.task_execution_log_request import TaskExecutionLogRequest + +class Q: + @staticmethod + def platforms() -> PlatformRequest: + return PlatformRequest(minimal=False) + + @staticmethod + def platforms_minimal() -> PlatformRequest: + return PlatformRequest(minimal=True) + + @staticmethod + def task_statuses() -> TaskStatusRequest: + return TaskStatusRequest(minimal=False) + + @staticmethod + def task_statuses_minimal() -> TaskStatusRequest: + return TaskStatusRequest(minimal=True) + + @staticmethod + def tasks() -> TaskRequest: + return TaskRequest(minimal=False) + + @staticmethod + def tasks_minimal() -> TaskRequest: + return TaskRequest(minimal=True) + + @staticmethod + def task_execution_logs() -> TaskExecutionLogRequest: + return TaskExecutionLogRequest(minimal=False) + + @staticmethod + def task_execution_logs_minimal() -> TaskExecutionLogRequest: + return TaskExecutionLogRequest(minimal=True) diff --git a/examples/task_board/generated/model.xml b/examples/task_board/generated/model.xml new file mode 100644 index 0000000..07e65d1 --- /dev/null +++ b/examples/task_board/generated/model.xml @@ -0,0 +1,55 @@ + + + + + <_value id="1001" name="Planned" code="PLANNED" color="#94A3B8" display_order="10" progress="0"/> + <_value id="1002" name="Ready" code="READY" color="#3B82F6" display_order="20" progress="25"/> + <_value id="1003" name="Executing" code="EXECUTING" color="#F59E0B" display_order="30" progress="50"/> + <_value id="1004" name="Verified" code="VERIFIED" color="#16A34A" display_order="40" progress="100"/> + + + + + + + + diff --git a/examples/task_board/generated/models/platform.py b/examples/task_board/generated/models/platform.py index 145dc4f..fc1ab44 100644 --- a/examples/task_board/generated/models/platform.py +++ b/examples/task_board/generated/models/platform.py @@ -1,58 +1,317 @@ 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 Platform: + _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 "founded" in kwargs and "founded" not in kwargs: + kwargs["founded"] = kwargs.pop("founded") + if "user_email" in kwargs and "userEmail" not in kwargs: + kwargs["userEmail"] = kwargs.pop("user_email") + 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()) self.id = kwargs.get("id") self.name = kwargs.get("name") self.founded = kwargs.get("founded") self.userEmail = kwargs.get("userEmail") self.version = kwargs.get("version") - for k, v in kwargs.items(): - if not hasattr(self, k): - setattr(self, k, v) + self._task_status_list = kwargs.get("task_status_list", []) + if "task_status_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("task_status_list") + self._task_list = kwargs.get("task_list", []) + if "task_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("task_list") + if self._task_status_list: + from models.task_status import TaskStatus + self._task_status_list = [ + item if isinstance(item, TaskStatus) else TaskStatus(**item) + for item in self._task_status_list + ] + if self._task_list: + from models.task import Task + self._task_list = [ + item if isinstance(item, Task) else Task(**item) + for item in self._task_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("Platform", 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._task_status_list: + child._teaql_attach_root(root) + for child in self._task_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, service): + async def save(self, context): + 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, "founded", None) is not None: - payload["founded"] = Value.I64(self.founded) - - if getattr(self, "userEmail", None) is not None: + if "founded" in self._loaded_fields: + payload["founded"] = Value.DateTime(self.founded) + if "userEmail" in self._loaded_fields: payload["user_email"] = Value.Text(self.userEmail) - - 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("Platform", payload) + elif action == "Update": + cmd = UpdateCommand("Platform", 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("Platform", 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 "founded" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("founded"), message="Mutation requires a fully loaded entity")]) + if "userEmail" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("user_email"), 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._task_status_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "platform", self) + child._loaded_fields.add("platform") + child._entity_root.set(child._teaql_entity_key(), "platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("task_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._task_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "platform", self) + child._loaded_fields.add("platform") + child._entity_root.set(child._teaql_entity_key(), "platform", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("task_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 - if self._action == "Create": - cmd = InsertCommand("Platform", payload) - elif self._action == "Update": - cmd = UpdateCommand("Platform", Value.from_any(getattr(self, "id", None))) - for k, v in payload.items(): - if k != "id": - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand("Platform", Value.from_any(getattr(self, "id", None))) + 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 - return await service.mutate(context, req) \ No newline at end of file + 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) + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for Platform" + ) + 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 "founded" in persisted: + self.founded = persisted["founded"] + self._loaded_fields.add("founded") + elif "founded" in persisted: + self.founded = persisted["founded"] + self._loaded_fields.add("founded") + if "user_email" in persisted: + self.userEmail = persisted["user_email"] + self._loaded_fields.add("userEmail") + elif "userEmail" in persisted: + self.userEmail = persisted["userEmail"] + self._loaded_fields.add("userEmail") + 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" + + cascade_relations = [] + cascade_relations.append(("task_status_list", self._task_status_list, "update_platform")) + cascade_relations.append(("task_list", self._task_list, "update_platform")) + if action != "Delete": + 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) + 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_founded(self, value): + self.founded = value + self._loaded_fields.add("founded") + self._entity_root.set(self._teaql_entity_key(), "founded", Value.from_any(value)) + return self + + def update_user_email(self, value): + self.userEmail = value + self._loaded_fields.add("userEmail") + self._entity_root.set(self._teaql_entity_key(), "user_email", 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 task_status_list(self) -> list: + self._loaded_fields.add("task_status_list") + return self._task_status_list + + def task_list(self) -> list: + self._loaded_fields.add("task_list") + return self._task_list diff --git a/examples/task_board/generated/models/task.py b/examples/task_board/generated/models/task.py index 4bd15df..c938894 100644 --- a/examples/task_board/generated/models/task.py +++ b/examples/task_board/generated/models/task.py @@ -1,58 +1,314 @@ 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.task_status import TaskStatus +from models.platform import Platform class Task: + _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 "status" in kwargs and "status" not in kwargs: + kwargs["status"] = kwargs.pop("status") + if "platform" in kwargs and "platform" not in kwargs: + kwargs["platform"] = kwargs.pop("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()) self.id = kwargs.get("id") self.name = kwargs.get("name") self.status = kwargs.get("status") self.platform = kwargs.get("platform") self.version = kwargs.get("version") - for k, v in kwargs.items(): - if not hasattr(self, k): - setattr(self, k, v) + if isinstance(self.status, dict): + self.status = TaskStatus(**self.status) + if isinstance(self.platform, dict): + self.platform = Platform(**self.platform) + self._task_execution_log_list = kwargs.get("task_execution_log_list", []) + if "task_execution_log_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("task_execution_log_list") + if self._task_execution_log_list: + from models.task_execution_log import TaskExecutionLog + self._task_execution_log_list = [ + item if isinstance(item, TaskExecutionLog) else TaskExecutionLog(**item) + for item in self._task_execution_log_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("Task", 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._task_execution_log_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, service): + async def save(self, context): + 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, "status", None) is not None: - payload["status"] = Value.I64(self.status) - - if getattr(self, "platform", None) is not None: - payload["platform"] = Value.I64(self.platform) - - if getattr(self, "version", None) is not None: + if "status" in self._loaded_fields: + payload["status"] = Value.Object(self.status) + if "platform" in self._loaded_fields: + payload["platform"] = Value.Object(self.platform) + 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("Task", payload) + elif action == "Update": + cmd = UpdateCommand("Task", 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("Task", 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 "status" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("status"), message="Mutation requires a fully loaded entity")]) + if "platform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("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._task_execution_log_list): + child._teaql_attach_root(self._entity_root) + setattr(child, "task", self) + child._loaded_fields.add("task") + child._entity_root.set(child._teaql_entity_key(), "task", Value.Object(self)) + child.audit_as(self._comment) + try: + child._teaql_preflight_graph(context) + except CheckException as error: + prefix = ObjectLocation().property("task_execution_log_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 - if self._action == "Create": - cmd = InsertCommand("Task", payload) - elif self._action == "Update": - cmd = UpdateCommand("Task", Value.from_any(getattr(self, "id", None))) - for k, v in payload.items(): - if k != "id": - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand("Task", Value.from_any(getattr(self, "id", None))) + 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 - return await service.mutate(context, req) \ No newline at end of file + 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) + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for Task" + ) + 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 "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 "platform" in persisted: + self.platform = persisted["platform"] + self._loaded_fields.add("platform") + elif "platform" in persisted: + self.platform = persisted["platform"] + self._loaded_fields.add("platform") + 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" + + cascade_relations = [] + cascade_relations.append(("task_execution_log_list", self._task_execution_log_list, "update_task")) + if action != "Delete": + 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) + 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_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_planned(self): + self.status = 1001 + 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_ready(self): + self.status = 1002 + 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_executing(self): + self.status = 1003 + 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_verified(self): + self.status = 1004 + self._loaded_fields.add("status") + self._entity_root.set(self._teaql_entity_key(), "status", Value.from_any(self.status)) + return self + + + def update_platform(self, value): + self.platform = getattr(value, "id", value) if value else None + self._loaded_fields.add("platform") + self._entity_root.set(self._teaql_entity_key(), "platform", Value.from_any(self.platform)) + return self + + def task_execution_log_list(self) -> list: + self._loaded_fields.add("task_execution_log_list") + return self._task_execution_log_list diff --git a/examples/task_board/generated/models/task_execution_log.py b/examples/task_board/generated/models/task_execution_log.py index e7553b8..15294c5 100644 --- a/examples/task_board/generated/models/task_execution_log.py +++ b/examples/task_board/generated/models/task_execution_log.py @@ -1,58 +1,260 @@ 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.task import Task class TaskExecutionLog: + _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 "task" in kwargs and "task" not in kwargs: + kwargs["task"] = kwargs.pop("task") + if "action" in kwargs and "action" not in kwargs: + kwargs["action"] = kwargs.pop("action") + if "detail" in kwargs and "detail" not in kwargs: + kwargs["detail"] = kwargs.pop("detail") + 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()) self.id = kwargs.get("id") self.task = kwargs.get("task") self.action = kwargs.get("action") self.detail = kwargs.get("detail") self.version = kwargs.get("version") - for k, v in kwargs.items(): - if not hasattr(self, k): - setattr(self, k, v) + if isinstance(self.task, dict): + self.task = Task(**self.task) + 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("TaskExecutionLog", 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, service): - payload = {} - if getattr(self, "id", None) is not None: - payload["id"] = Value.I64(self.id) + async def save(self, context): + return await context.execute_graph_save(lambda: self._teaql_preflight_and_save(context)) - if getattr(self, "task", None) is not None: - payload["task"] = Value.I64(self.task) + async def _teaql_preflight_and_save(self, context): + self._teaql_preflight_graph(context) + return await self._teaql_save_within_graph(context) - if getattr(self, "action", None) is not None: + def _teaql_build_command(self): + payload = {} + if "id" in self._loaded_fields: + payload["id"] = Value.I64(self.id) + if "task" in self._loaded_fields: + payload["task"] = Value.Object(self.task) + if "action" in self._loaded_fields: payload["action"] = Value.Text(self.action) - - if getattr(self, "detail", None) is not None: + if "detail" in self._loaded_fields: payload["detail"] = Value.Text(self.detail) - - 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("TaskExecutionLog", payload) + elif action == "Update": + cmd = UpdateCommand("TaskExecutionLog", 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("TaskExecutionLog", 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 "task" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("task"), message="Mutation requires a fully loaded entity")]) + if "action" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("action"), message="Mutation requires a fully loaded entity")]) + if "detail" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("detail"), 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) - if self._action == "Create": - cmd = InsertCommand("TaskExecutionLog", payload) - elif self._action == "Update": - cmd = UpdateCommand("TaskExecutionLog", Value.from_any(getattr(self, "id", None))) - for k, v in payload.items(): - if k != "id": - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand("TaskExecutionLog", Value.from_any(getattr(self, "id", None))) + 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 - return await service.mutate(context, req) \ No newline at end of file + 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) + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for TaskExecutionLog" + ) + 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 "task" in persisted: + self.task = persisted["task"] + self._loaded_fields.add("task") + elif "task" in persisted: + self.task = persisted["task"] + self._loaded_fields.add("task") + if "action" in persisted: + self.action = persisted["action"] + self._loaded_fields.add("action") + elif "action" in persisted: + self.action = persisted["action"] + self._loaded_fields.add("action") + if "detail" in persisted: + self.detail = persisted["detail"] + self._loaded_fields.add("detail") + elif "detail" in persisted: + self.detail = persisted["detail"] + self._loaded_fields.add("detail") + 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" + + cascade_relations = [] + if action != "Delete": + 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) + 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_action(self, value): + self.action = value + self._loaded_fields.add("action") + self._entity_root.set(self._teaql_entity_key(), "action", Value.from_any(value)) + return self + + def update_detail(self, value): + self.detail = value + self._loaded_fields.add("detail") + self._entity_root.set(self._teaql_entity_key(), "detail", 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_task(self, value): + self.task = getattr(value, "id", value) if value else None + self._loaded_fields.add("task") + self._entity_root.set(self._teaql_entity_key(), "task", Value.from_any(self.task)) + return self diff --git a/examples/task_board/generated/models/task_status.py b/examples/task_board/generated/models/task_status.py index c2482d4..047cad1 100644 --- a/examples/task_board/generated/models/task_status.py +++ b/examples/task_board/generated/models/task_status.py @@ -1,10 +1,46 @@ 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.platform import Platform class TaskStatus: + _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 "progress" in kwargs and "progress" not in kwargs: + kwargs["progress"] = kwargs.pop("progress") + if "platform" in kwargs and "platform" not in kwargs: + kwargs["platform"] = kwargs.pop("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()) self.id = kwargs.get("id") self.name = kwargs.get("name") self.code = kwargs.get("code") @@ -13,58 +49,299 @@ def __init__(self, **kwargs): self.progress = kwargs.get("progress") self.platform = kwargs.get("platform") self.version = kwargs.get("version") - for k, v in kwargs.items(): - if not hasattr(self, k): - setattr(self, k, v) + if isinstance(self.platform, dict): + self.platform = Platform(**self.platform) + self._task_list = kwargs.get("task_list", []) + if "task_list" in kwargs or kwargs.get("id") is None: + self._loaded_fields.add("task_list") + if self._task_list: + from models.task import Task + self._task_list = [ + item if isinstance(item, Task) else Task(**item) + for item in self._task_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("TaskStatus", 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._task_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, service): + async def save(self, context): + 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.I64(self.displayOrder) - - if getattr(self, "progress", None) is not None: - payload["progress"] = Value.I64(self.progress) - - if getattr(self, "platform", None) is not None: - payload["platform"] = Value.I64(self.platform) - - if getattr(self, "version", None) is not None: + if "displayOrder" in self._loaded_fields: + payload["display_order"] = Value.Decimal(self.displayOrder) + if "progress" in self._loaded_fields: + payload["progress"] = Value.Decimal(self.progress) + if "platform" in self._loaded_fields: + payload["platform"] = Value.Object(self.platform) + 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("TaskStatus", payload) + elif action == "Update": + cmd = UpdateCommand("TaskStatus", 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("TaskStatus", 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 "progress" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("progress"), message="Mutation requires a fully loaded entity")]) + if "platform" not in self._loaded_fields: + raise CheckException([CheckResult("invalid_type", ObjectLocation().property("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._task_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("task_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 - if self._action == "Create": - cmd = InsertCommand("TaskStatus", payload) - elif self._action == "Update": - cmd = UpdateCommand("TaskStatus", Value.from_any(getattr(self, "id", None))) - for k, v in payload.items(): - if k != "id": - cmd.value(k, v) - elif self._action == "Delete": - cmd = DeleteCommand("TaskStatus", Value.from_any(getattr(self, "id", None))) + 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 - return await service.mutate(context, req) \ No newline at end of file + 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) + persisted = result.persisted_record + if persisted is None: + raise RuntimeError( + "Mutation provider did not return authoritative persisted state for TaskStatus" + ) + 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 "progress" in persisted: + self.progress = persisted["progress"] + self._loaded_fields.add("progress") + elif "progress" in persisted: + self.progress = persisted["progress"] + self._loaded_fields.add("progress") + if "platform" in persisted: + self.platform = persisted["platform"] + self._loaded_fields.add("platform") + elif "platform" in persisted: + self.platform = persisted["platform"] + self._loaded_fields.add("platform") + 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" + + cascade_relations = [] + cascade_relations.append(("task_list", self._task_list, "update_status")) + if action != "Delete": + 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) + 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_progress(self, value): + self.progress = value + self._loaded_fields.add("progress") + self._entity_root.set(self._teaql_entity_key(), "progress", 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_platform(self, value): + self.platform = getattr(value, "id", value) if value else None + self._loaded_fields.add("platform") + self._entity_root.set(self._teaql_entity_key(), "platform", Value.from_any(self.platform)) + return self + + def task_list(self) -> list: + self._loaded_fields.add("task_list") + return self._task_list diff --git a/examples/task_board/generated/pyproject.toml b/examples/task_board/generated/pyproject.toml new file mode 100644 index 0000000..bee20bc --- /dev/null +++ b/examples/task_board/generated/pyproject.toml @@ -0,0 +1,16 @@ +[project] +name = "robot-kanban-service-lib" +version = "1.0.0" +description = "Generated python library" +dependencies = ["teaql==0.2.7", "aiosqlite>=0.22.1"] + +[tool.setuptools] +py-modules = ["Q", "E"] + +[tool.setuptools.packages.find] +where = ["."] +include = ["models*", "requests*"] + +[build-system] +requires = ["setuptools>=42"] +build-backend = "setuptools.build_meta" \ No newline at end of file diff --git a/examples/task_board/generated/requests/platform_request.py b/examples/task_board/generated/requests/platform_request.py index 3d453dd..c47f851 100644 --- a/examples/task_board/generated/requests/platform_request.py +++ b/examples/task_board/generated/requests/platform_request.py @@ -1,46 +1,416 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, 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.platform import Platform +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class PlatformRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("Platform") + 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) + self._comment = c return self def purpose(self, p: str): + self.query.purpose(p) + self._purpose = p + return ExecutablePlatformRequest(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 + + def offset(self, n: int): + self.query.offset(n) + return self + + def with_deleted_rows(self): + self.query.with_deleted_rows() return self + def deleted_rows_only(self): + self.query.deleted_rows_only() + return self + + def select_self_fields(self): + self.query.project("id", "name", "founded", "user_email", "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_founded(self): + self.query.project("founded") + return self + + def select_user_email(self): + self.query.project("user_email") + 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_founded_is(self, val): self.query.and_filter(eq("founded", val)) return self + def with_founded_is_not(self, val): + self.query.and_filter(ne("founded", val)) + return self + + def with_founded_in(self, *vals): + self.query.and_filter(in_list("founded", list(vals))) + return self + + def with_founded_not_in(self, *vals): + self.query.and_filter(not_in_list("founded", list(vals))) + return self + + def with_founded_greater_than(self, val): + self.query.and_filter(gt("founded", val)) + return self + + def with_founded_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("founded", val)) + return self + + def with_founded_less_than(self, val): + self.query.and_filter(lt("founded", val)) + return self + + def with_founded_less_than_or_equal_to(self, val): + self.query.and_filter(lte("founded", val)) + return self + + def with_founded_between(self, lower, upper): + self.query.and_filter(between(column("founded"), value(lower), value(upper))) + return self + + def with_founded_is_known(self): + self.query.and_filter(is_not_null(column("founded"))) + return self + + def with_founded_is_unknown(self): + self.query.and_filter(is_null(column("founded"))) + return self + def with_user_email_containing(self, val: str): self.query.and_filter(contain("user_email", val)) return self + def with_user_email_not_containing(self, val: str): + self.query.and_filter(not_contain("user_email", val)) + return self + + def with_user_email_starting_with(self, val: str): + self.query.and_filter(begin_with("user_email", val)) + return self + + def with_user_email_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("user_email", val)) + return self + + def with_user_email_ending_with(self, val: str): + self.query.and_filter(end_with("user_email", val)) + return self + + def with_user_email_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("user_email", val)) + return self + + def with_user_email_sounding_like(self, val: str): + self.query.and_filter(sound_like("user_email", val)) + return self + def with_user_email_is(self, val: str): self.query.and_filter(eq("user_email", val)) return self + def with_user_email_is_not(self, val): + self.query.and_filter(ne("user_email", val)) + return self + + def with_user_email_in(self, *vals): + self.query.and_filter(in_list("user_email", list(vals))) + return self + + def with_user_email_not_in(self, *vals): + self.query.and_filter(not_in_list("user_email", list(vals))) + return self + + def with_user_email_greater_than(self, val): + self.query.and_filter(gt("user_email", val)) + return self + + def with_user_email_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("user_email", val)) + return self + + def with_user_email_less_than(self, val): + self.query.and_filter(lt("user_email", val)) + return self + + def with_user_email_less_than_or_equal_to(self, val): + self.query.and_filter(lte("user_email", val)) + return self + + def with_user_email_between(self, lower, upper): + self.query.and_filter(between(column("user_email"), value(lower), value(upper))) + return self + + def with_user_email_is_known(self): + self.query.and_filter(is_not_null(column("user_email"))) + return self + + def with_user_email_is_unknown(self): + self.query.and_filter(is_null(column("user_email"))) + 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 + + def order_by_id_descending(self): + self.query.order_by("id", "desc") + return self + + def order_by_name_ascending(self): + self.query.order_by("name", "asc") + return self + + def order_by_name_descending(self): + self.query.order_by("name", "desc") + return self + + def order_by_founded_ascending(self): + self.query.order_by("founded", "asc") + return self + + def order_by_founded_descending(self): + self.query.order_by("founded", "desc") + return self + + def order_by_user_email_ascending(self): + self.query.order_by("user_email", "asc") + return self + + def order_by_user_email_descending(self): + self.query.order_by("user_email", "desc") + return self + + def order_by_version_ascending(self): + self.query.order_by("version", "asc") + return self + + def order_by_version_descending(self): + self.query.order_by("version", "desc") + return self + def count(self): self.query.count_field("id", "count") @@ -55,45 +425,367 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self - def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self - def group_by_founded(self): self.query.group_by("founded") return self def group_by_founded_as(self, ret_name: str): - self.query.group_by("founded") + self.query.group_by("founded") return self - def group_by_user_email(self): self.query.group_by("user_email") return self def group_by_user_email_as(self, ret_name: str): - self.query.group_by("user_email") + self.query.group_by("user_email") return self - def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") + return self + def select_task_status_list(self): + from requests.task_status_request import TaskStatusRequest + return self.select_task_status_list_with(TaskStatusRequest()) + + def select_task_status_list_with(self, child_request): + self.query.relation_query("task_status_list", child_request.query) return self + def select_task_list(self): + from requests.task_request import TaskRequest + return self.select_task_list_with(TaskRequest()) + + def select_task_list_with(self, child_request): + self.query.relation_query("task_list", child_request.query) + return self + def have_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.with_task_status_list_matching(TaskStatusRequest()) + + def have_no_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.without_task_status_list_matching(TaskStatusRequest()) + + def with_task_status_list_matching(self, child_request): + child_request.query.projection = ["platform"] + self.query.and_filter(in_subquery(column("id"), "TaskStatus", child_request.query)) + return self + + def without_task_status_list_matching(self, child_request): + child_request.query.projection = ["platform"] + self.query.and_filter(not_in_subquery(column("id"), "TaskStatus", child_request.query)) + return self + def have_tasks(self): + from requests.task_request import TaskRequest + return self.with_task_list_matching(TaskRequest()) + + def have_no_tasks(self): + from requests.task_request import TaskRequest + return self.without_task_list_matching(TaskRequest()) + + def with_task_list_matching(self, child_request): + child_request.query.projection = ["platform"] + self.query.and_filter(in_subquery(column("id"), "Task", child_request.query)) + return self + + def without_task_list_matching(self, child_request): + child_request.query.projection = ["platform"] + self.query.and_filter(not_in_subquery(column("id"), "Task", child_request.query)) + return self + def count_task_statuses(self): + return self.count_task_statuses_as("count_task_statuses") + + def count_task_statuses_as(self, alias: str): + from requests.task_status_request import TaskStatusRequest + return self.count_task_statuses_with(alias, TaskStatusRequest()) + + def count_task_statuses_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + + def min_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.min_display_order_of_task_statuses_as( + "min_display_order_of_task_statuses", TaskStatusRequest()) + + def min_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.min("display_order", "min_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def max_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.max_display_order_of_task_statuses_as( + "max_display_order_of_task_statuses", TaskStatusRequest()) + + def max_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.max("display_order", "max_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def sum_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.sum_display_order_of_task_statuses_as( + "sum_display_order_of_task_statuses", TaskStatusRequest()) + + def sum_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.sum("display_order", "sum_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def avg_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.avg_display_order_of_task_statuses_as( + "avg_display_order_of_task_statuses", TaskStatusRequest()) + + def avg_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.avg("display_order", "avg_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def standardDeviation_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.standardDeviation_display_order_of_task_statuses_as( + "standardDeviation_display_order_of_task_statuses", TaskStatusRequest()) + + def standardDeviation_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.standardDeviation("display_order", "standardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def squareRootOfPopulationStandardDeviation_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.squareRootOfPopulationStandardDeviation_display_order_of_task_statuses_as( + "squareRootOfPopulationStandardDeviation_display_order_of_task_statuses", TaskStatusRequest()) + + def squareRootOfPopulationStandardDeviation_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.squareRootOfPopulationStandardDeviation("display_order", "squareRootOfPopulationStandardDeviation_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def sampleVariance_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.sampleVariance_display_order_of_task_statuses_as( + "sampleVariance_display_order_of_task_statuses", TaskStatusRequest()) + + def sampleVariance_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.sampleVariance("display_order", "sampleVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def samplePopulationVariance_display_order_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.samplePopulationVariance_display_order_of_task_statuses_as( + "samplePopulationVariance_display_order_of_task_statuses", TaskStatusRequest()) + + def samplePopulationVariance_display_order_of_task_statuses_as(self, alias: str, child_request): + child_request.query.samplePopulationVariance("display_order", "samplePopulationVariance_display_order") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def min_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.min_progress_of_task_statuses_as( + "min_progress_of_task_statuses", TaskStatusRequest()) + + def min_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.min("progress", "min_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def max_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.max_progress_of_task_statuses_as( + "max_progress_of_task_statuses", TaskStatusRequest()) + + def max_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.max("progress", "max_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def sum_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.sum_progress_of_task_statuses_as( + "sum_progress_of_task_statuses", TaskStatusRequest()) + + def sum_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.sum("progress", "sum_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def avg_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.avg_progress_of_task_statuses_as( + "avg_progress_of_task_statuses", TaskStatusRequest()) + + def avg_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.avg("progress", "avg_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def standardDeviation_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.standardDeviation_progress_of_task_statuses_as( + "standardDeviation_progress_of_task_statuses", TaskStatusRequest()) + + def standardDeviation_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.standardDeviation("progress", "standardDeviation_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def squareRootOfPopulationStandardDeviation_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.squareRootOfPopulationStandardDeviation_progress_of_task_statuses_as( + "squareRootOfPopulationStandardDeviation_progress_of_task_statuses", TaskStatusRequest()) + + def squareRootOfPopulationStandardDeviation_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.squareRootOfPopulationStandardDeviation("progress", "squareRootOfPopulationStandardDeviation_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def sampleVariance_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.sampleVariance_progress_of_task_statuses_as( + "sampleVariance_progress_of_task_statuses", TaskStatusRequest()) + + def sampleVariance_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.sampleVariance("progress", "sampleVariance_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def samplePopulationVariance_progress_of_task_statuses(self): + from requests.task_status_request import TaskStatusRequest + return self.samplePopulationVariance_progress_of_task_statuses_as( + "samplePopulationVariance_progress_of_task_statuses", TaskStatusRequest()) + + def samplePopulationVariance_progress_of_task_statuses_as(self, alias: str, child_request): + child_request.query.samplePopulationVariance("progress", "samplePopulationVariance_progress") + self.query.relation_aggregates.append( + RelationAggregate("task_status_list", alias, child_request.query, True) + ) + return self + def count_tasks(self): + return self.count_tasks_as("count_tasks") + + def count_tasks_as(self, alias: str): + from requests.task_request import TaskRequest + return self.count_tasks_with(alias, TaskRequest()) + + def count_tasks_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregates.append( + RelationAggregate("task_list", alias, child_request.query, True) + ) + return self + + + +class ExecutablePlatformRequest: + def __init__(self, request): + self._request = request + + def comment(self, c: str): + self._request.comment(c) + return self + + def new_entity(self, context) -> Platform: + 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("Platform", Platform()) + if not isinstance(entity, Platform): + raise TypeError("entity initializer returned an incompatible Platform") + return entity + + async def execute_for_result(self, context): + self = self._request + 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(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) + 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[Platform]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (Platform(_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[Platform]: + 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, _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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, _comment=request._comment, _purpose=request._purpose)) + query_root = EntityRoot() + data = SmartList(Platform(_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_list(self, context, service): - req = QueryRequest(self.query) - res = await service.query(context, req) + async def execute_for_one(self, context): + self._request.limit(1) + entities = await self.execute_for_list(context) + return entities[0] if entities else None - result = {"data": res.rows} - return result \ No newline at end of file + 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(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): + for row in chunk.rows: + yield Platform(_entity_root=query_root, **row) diff --git a/examples/task_board/generated/requests/task_execution_log_request.py b/examples/task_board/generated/requests/task_execution_log_request.py index 4cb31c9..1cbb08a 100644 --- a/examples/task_board/generated/requests/task_execution_log_request.py +++ b/examples/task_board/generated/requests/task_execution_log_request.py @@ -1,43 +1,387 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, 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.task_execution_log import TaskExecutionLog +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class TaskExecutionLogRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("TaskExecutionLog") + 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) + self._comment = c return self def purpose(self, p: str): + self.query.purpose(p) + self._purpose = p + return ExecutableTaskExecutionLogRequest(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 + + def offset(self, n: int): + self.query.offset(n) + return self + + def with_deleted_rows(self): + self.query.with_deleted_rows() + return self + + def deleted_rows_only(self): + self.query.deleted_rows_only() + return self + + def select_self_fields(self): + self.query.project("id", "task", "action", "detail", "version") + return self + + def select_id(self): + self.query.project("id") + return self + + + def select_action(self): + self.query.project("action") + return self + + def select_detail(self): + self.query.project("detail") + return self + + def select_version(self): + self.query.project("version") + return self + + def select_task_with(self, child_request): + self.query.project("task") + self.query.relation_query("task", child_request.query) + return self + def with_task_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("task"), "Task", child_request.query)) + return self + + def without_task_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("task"), "Task", child_request.query)) + return self + + def have_task(self): + self.query.and_filter(is_not_null(column("task"))) + return self + + def have_no_task(self): + self.query.and_filter(is_null(column("task"))) 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_task(self, val): + self.query.and_filter(eq("task", val)) + return self def with_action_containing(self, val: str): self.query.and_filter(contain("action", val)) return self + def with_action_not_containing(self, val: str): + self.query.and_filter(not_contain("action", val)) + return self + + def with_action_starting_with(self, val: str): + self.query.and_filter(begin_with("action", val)) + return self + + def with_action_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("action", val)) + return self + + def with_action_ending_with(self, val: str): + self.query.and_filter(end_with("action", val)) + return self + + def with_action_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("action", val)) + return self + + def with_action_sounding_like(self, val: str): + self.query.and_filter(sound_like("action", val)) + return self + def with_action_is(self, val: str): self.query.and_filter(eq("action", val)) return self + def with_action_is_not(self, val): + self.query.and_filter(ne("action", val)) + return self + + def with_action_in(self, *vals): + self.query.and_filter(in_list("action", list(vals))) + return self + + def with_action_not_in(self, *vals): + self.query.and_filter(not_in_list("action", list(vals))) + return self + + def with_action_greater_than(self, val): + self.query.and_filter(gt("action", val)) + return self + + def with_action_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("action", val)) + return self + + def with_action_less_than(self, val): + self.query.and_filter(lt("action", val)) + return self + + def with_action_less_than_or_equal_to(self, val): + self.query.and_filter(lte("action", val)) + return self + + def with_action_between(self, lower, upper): + self.query.and_filter(between(column("action"), value(lower), value(upper))) + return self + + def with_action_is_known(self): + self.query.and_filter(is_not_null(column("action"))) + return self + + def with_action_is_unknown(self): + self.query.and_filter(is_null(column("action"))) + return self def with_detail_containing(self, val: str): self.query.and_filter(contain("detail", val)) return self + def with_detail_not_containing(self, val: str): + self.query.and_filter(not_contain("detail", val)) + return self + + def with_detail_starting_with(self, val: str): + self.query.and_filter(begin_with("detail", val)) + return self + + def with_detail_not_starting_with(self, val: str): + self.query.and_filter(not_begin_with("detail", val)) + return self + + def with_detail_ending_with(self, val: str): + self.query.and_filter(end_with("detail", val)) + return self + + def with_detail_not_ending_with(self, val: str): + self.query.and_filter(not_end_with("detail", val)) + return self + + def with_detail_sounding_like(self, val: str): + self.query.and_filter(sound_like("detail", val)) + return self + def with_detail_is(self, val: str): self.query.and_filter(eq("detail", val)) return self + def with_detail_is_not(self, val): + self.query.and_filter(ne("detail", val)) + return self + + def with_detail_in(self, *vals): + self.query.and_filter(in_list("detail", list(vals))) + return self + + def with_detail_not_in(self, *vals): + self.query.and_filter(not_in_list("detail", list(vals))) + return self + + def with_detail_greater_than(self, val): + self.query.and_filter(gt("detail", val)) + return self + + def with_detail_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("detail", val)) + return self + + def with_detail_less_than(self, val): + self.query.and_filter(lt("detail", val)) + return self + + def with_detail_less_than_or_equal_to(self, val): + self.query.and_filter(lte("detail", val)) + return self + + def with_detail_between(self, lower, upper): + self.query.and_filter(between(column("detail"), value(lower), value(upper))) + return self + + def with_detail_is_known(self): + self.query.and_filter(is_not_null(column("detail"))) + return self + + def with_detail_is_unknown(self): + self.query.and_filter(is_null(column("detail"))) + 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 + + def order_by_id_descending(self): + self.query.order_by("id", "desc") + return self + + + def order_by_action_ascending(self): + self.query.order_by("action", "asc") + return self + + def order_by_action_descending(self): + self.query.order_by("action", "desc") + return self + + def order_by_detail_ascending(self): + self.query.order_by("detail", "asc") + return self + + def order_by_detail_descending(self): + self.query.order_by("detail", "desc") + return self + + def order_by_version_ascending(self): + self.query.order_by("version", "asc") + return self + + def order_by_version_descending(self): + self.query.order_by("version", "desc") + return self + def count(self): self.query.count_field("id", "count") @@ -52,38 +396,119 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") + return self + def group_by_task(self): + self.query.group_by("task") return self - + def group_by_task_as(self, ret_name: str): + self.query.group_by("task") + return self def group_by_action(self): self.query.group_by("action") return self def group_by_action_as(self, ret_name: str): - self.query.group_by("action") + self.query.group_by("action") return self - def group_by_detail(self): self.query.group_by("detail") return self def group_by_detail_as(self, ret_name: str): - self.query.group_by("detail") + self.query.group_by("detail") return self - def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") + return self + def facet_by_task_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "task", request.query, include_all_facets) return self - async def execute_for_list(self, context, service): - req = QueryRequest(self.query) - res = await service.query(context, req) +class ExecutableTaskExecutionLogRequest: + def __init__(self, request): + self._request = request + + def comment(self, c: str): + self._request.comment(c) + return self + + def new_entity(self, context) -> TaskExecutionLog: + 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("TaskExecutionLog", TaskExecutionLog()) + if not isinstance(entity, TaskExecutionLog): + raise TypeError("entity initializer returned an incompatible TaskExecutionLog") + return entity + + async def execute_for_result(self, context): + self = self._request + 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(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) + 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[TaskExecutionLog]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (TaskExecutionLog(_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[TaskExecutionLog]: + 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, _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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, _comment=request._comment, _purpose=request._purpose)) + query_root = EntityRoot() + data = SmartList(TaskExecutionLog(_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) + entities = await self.execute_for_list(context) + return entities[0] if entities else None - result = {"data": res.rows} - return result \ No newline at end of file + 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(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): + for row in chunk.rows: + yield TaskExecutionLog(_entity_root=query_root, **row) diff --git a/examples/task_board/generated/requests/task_request.py b/examples/task_board/generated/requests/task_request.py index 4a95936..eaa6d57 100644 --- a/examples/task_board/generated/requests/task_request.py +++ b/examples/task_board/generated/requests/task_request.py @@ -1,36 +1,343 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, 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.task import Task +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class TaskRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("Task") + 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) + self._comment = c return self def purpose(self, p: str): + self.query.purpose(p) + self._purpose = p + return ExecutableTaskRequest(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 + + def offset(self, n: int): + self.query.offset(n) + return self + + def with_deleted_rows(self): + self.query.with_deleted_rows() + return self + + def deleted_rows_only(self): + self.query.deleted_rows_only() + return self + + def select_self_fields(self): + self.query.project("id", "name", "status", "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_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_platform_with(self, child_request): + self.query.project("platform") + self.query.relation_query("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"), "TaskStatus", 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"), "TaskStatus", 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_platform_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("platform"), "Platform", child_request.query)) + return self + + def without_platform_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_request.query)) + return self + + def have_platform(self): + self.query.and_filter(is_not_null(column("platform"))) + return self + + def have_no_platform(self): + self.query.and_filter(is_null(column("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 filter_by_status(self, val): + self.query.and_filter(eq("status", val)) + return self + def with_status_is_planned(self): + self.query.and_filter(eq("status", 1001)) + return self + def with_status_is_ready(self): + self.query.and_filter(eq("status", 1002)) + return self + def with_status_is_executing(self): + self.query.and_filter(eq("status", 1003)) + return self + def with_status_is_verified(self): + self.query.and_filter(eq("status", 1004)) + return self + + def filter_by_platform(self, val): + self.query.and_filter(eq("platform", val)) + 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 + + def order_by_id_descending(self): + self.query.order_by("id", "desc") + return self + + def order_by_name_ascending(self): + self.query.order_by("name", "asc") + return self + + def order_by_name_descending(self): + self.query.order_by("name", "desc") + return self + + + + def order_by_version_ascending(self): + self.query.order_by("version", "asc") + return self + + def order_by_version_descending(self): + self.query.order_by("version", "desc") + return self + def count(self): self.query.count_field("id", "count") @@ -45,31 +352,163 @@ def group_by_id(self): return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self - def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") + return self + def group_by_status(self): + self.query.group_by("status") return self + def group_by_status_as(self, ret_name: str): + self.query.group_by("status") + return self + def group_by_platform(self): + self.query.group_by("platform") + return self - + def group_by_platform_as(self, ret_name: str): + self.query.group_by("platform") + return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") + return self + def select_task_execution_log_list(self): + from requests.task_execution_log_request import TaskExecutionLogRequest + return self.select_task_execution_log_list_with(TaskExecutionLogRequest()) + + def select_task_execution_log_list_with(self, child_request): + self.query.relation_query("task_execution_log_list", child_request.query) + return self + def have_task_execution_logs(self): + from requests.task_execution_log_request import TaskExecutionLogRequest + return self.with_task_execution_log_list_matching(TaskExecutionLogRequest()) + + def have_no_task_execution_logs(self): + from requests.task_execution_log_request import TaskExecutionLogRequest + return self.without_task_execution_log_list_matching(TaskExecutionLogRequest()) + + def with_task_execution_log_list_matching(self, child_request): + child_request.query.projection = ["task"] + self.query.and_filter(in_subquery(column("id"), "TaskExecutionLog", child_request.query)) + return self + + def without_task_execution_log_list_matching(self, child_request): + child_request.query.projection = ["task"] + self.query.and_filter(not_in_subquery(column("id"), "TaskExecutionLog", child_request.query)) return self + def count_task_execution_logs(self): + return self.count_task_execution_logs_as("count_task_execution_logs") + + def count_task_execution_logs_as(self, alias: str): + from requests.task_execution_log_request import TaskExecutionLogRequest + return self.count_task_execution_logs_with(alias, TaskExecutionLogRequest()) + + def count_task_execution_logs_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregates.append( + RelationAggregate("task_execution_log_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_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "platform", request.query, include_all_facets) + return self + + +class ExecutableTaskRequest: + def __init__(self, request): + self._request = request + + def comment(self, c: str): + self._request.comment(c) + return self + + def new_entity(self, context) -> Task: + 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("Task", Task()) + if not isinstance(entity, Task): + raise TypeError("entity initializer returned an incompatible Task") + return entity + + async def execute_for_result(self, context): + self = self._request + 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(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) + 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[Task]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (Task(_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[Task]: + 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, _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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, _comment=request._comment, _purpose=request._purpose)) + query_root = EntityRoot() + data = SmartList(Task(_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_list(self, context, service): - req = QueryRequest(self.query) - res = await service.query(context, req) + async def execute_for_one(self, context): + self._request.limit(1) + entities = await self.execute_for_list(context) + return entities[0] if entities else None - result = {"data": res.rows} - return result \ No newline at end of file + 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(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): + for row in chunk.rows: + yield Task(_entity_root=query_root, **row) diff --git a/examples/task_board/generated/requests/task_status_request.py b/examples/task_board/generated/requests/task_status_request.py index ffa73b9..737eb50 100644 --- a/examples/task_board/generated/requests/task_status_request.py +++ b/examples/task_board/generated/requests/task_status_request.py @@ -1,59 +1,582 @@ -from teaql.core.query import SelectQuery +from teaql.core.query import RelationAggregate, 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.task_status import TaskStatus +from typing import Protocol + +class QuerySelection(Protocol): + query: SelectQuery class TaskStatusRequest: - def __init__(self): + def __init__(self, minimal=False): self.query = SelectQuery("TaskStatus") + 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) + self._comment = c return self def purpose(self, p: str): + self.query.purpose(p) + self._purpose = p + return ExecutableTaskStatusRequest(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 + + def offset(self, n: int): + self.query.offset(n) + return self + + def with_deleted_rows(self): + self.query.with_deleted_rows() + return self + + def deleted_rows_only(self): + self.query.deleted_rows_only() + return self + + def select_self_fields(self): + self.query.project("id", "name", "code", "color", "display_order", "progress", "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_progress(self): + self.query.project("progress") + return self + + + def select_version(self): + self.query.project("version") + return self + + def select_platform_with(self, child_request): + self.query.project("platform") + self.query.relation_query("platform", child_request.query) + return self + def with_platform_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(in_subquery(column("platform"), "Platform", child_request.query)) + return self + + def without_platform_matching(self, child_request): + child_request.query.projection = ["id"] + self.query.and_filter(not_in_subquery(column("platform"), "Platform", child_request.query)) + return self + + def have_platform(self): + self.query.and_filter(is_not_null(column("platform"))) + return self + + def have_no_platform(self): + self.query.and_filter(is_null(column("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 with_progress_is(self, val): self.query.and_filter(eq("progress", val)) return self + def with_progress_is_not(self, val): + self.query.and_filter(ne("progress", val)) + return self + + def with_progress_in(self, *vals): + self.query.and_filter(in_list("progress", list(vals))) + return self + + def with_progress_not_in(self, *vals): + self.query.and_filter(not_in_list("progress", list(vals))) + return self + + def with_progress_greater_than(self, val): + self.query.and_filter(gt("progress", val)) + return self + + def with_progress_greater_than_or_equal_to(self, val): + self.query.and_filter(gte("progress", val)) + return self + + def with_progress_less_than(self, val): + self.query.and_filter(lt("progress", val)) + return self + + def with_progress_less_than_or_equal_to(self, val): + self.query.and_filter(lte("progress", val)) + return self + + def with_progress_between(self, lower, upper): + self.query.and_filter(between(column("progress"), value(lower), value(upper))) + return self + + def with_progress_is_known(self): + self.query.and_filter(is_not_null(column("progress"))) + return self + + def with_progress_is_unknown(self): + self.query.and_filter(is_null(column("progress"))) + return self + + def filter_by_platform(self, val): + self.query.and_filter(eq("platform", val)) + 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 + + def order_by_id_descending(self): + self.query.order_by("id", "desc") + return self + + def order_by_name_ascending(self): + self.query.order_by("name", "asc") + return self + + def order_by_name_descending(self): + self.query.order_by("name", "desc") + return self + + def order_by_code_ascending(self): + self.query.order_by("code", "asc") + return self + + def order_by_code_descending(self): + self.query.order_by("code", "desc") + return self + + def order_by_color_ascending(self): + self.query.order_by("color", "asc") + return self + + def order_by_color_descending(self): + self.query.order_by("color", "desc") + return self + + def order_by_display_order_ascending(self): + self.query.order_by("display_order", "asc") + return self + + def order_by_display_order_descending(self): + self.query.order_by("display_order", "desc") + return self + + def order_by_progress_ascending(self): + self.query.order_by("progress", "asc") + return self + + def order_by_progress_descending(self): + self.query.order_by("progress", "desc") + return self + + + def order_by_version_ascending(self): + self.query.order_by("version", "asc") + return self + + def order_by_version_descending(self): + self.query.order_by("version", "desc") + return self + def count(self): self.query.count_field("id", "count") @@ -67,159 +590,276 @@ def min_display_order(self): return self.min_display_order_as("minOfDisplayOrder") def min_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.min("display_order", ret_name) return self def max_display_order(self): return self.max_display_order_as("maxOfDisplayOrder") def max_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.max("display_order", ret_name) return self def sum_display_order(self): return self.sum_display_order_as("sumOfDisplayOrder") def sum_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.sum("display_order", ret_name) return self def avg_display_order(self): return self.avg_display_order_as("avgOfDisplayOrder") def avg_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.avg("display_order", ret_name) return self def standardDeviation_display_order(self): return self.standardDeviation_display_order_as("standardDeviationOfDisplayOrder") def standardDeviation_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.standardDeviation("display_order", ret_name) return self def squareRootOfPopulationStandardDeviation_display_order(self): return self.squareRootOfPopulationStandardDeviation_display_order_as("squareRootOfPopulationStandardDeviationOfDisplayOrder") def squareRootOfPopulationStandardDeviation_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.squareRootOfPopulationStandardDeviation("display_order", ret_name) return self def sampleVariance_display_order(self): return self.sampleVariance_display_order_as("sampleVarianceOfDisplayOrder") def sampleVariance_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.sampleVariance("display_order", ret_name) return self def samplePopulationVariance_display_order(self): return self.samplePopulationVariance_display_order_as("samplePopulationVarianceOfDisplayOrder") def samplePopulationVariance_display_order_as(self, ret_name: str): - self.query.("display_order", ret_name) + self.query.samplePopulationVariance("display_order", ret_name) return self def min_progress(self): return self.min_progress_as("minOfProgress") def min_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.min("progress", ret_name) return self def max_progress(self): return self.max_progress_as("maxOfProgress") def max_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.max("progress", ret_name) return self def sum_progress(self): return self.sum_progress_as("sumOfProgress") def sum_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.sum("progress", ret_name) return self def avg_progress(self): return self.avg_progress_as("avgOfProgress") def avg_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.avg("progress", ret_name) return self def standardDeviation_progress(self): return self.standardDeviation_progress_as("standardDeviationOfProgress") def standardDeviation_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.standardDeviation("progress", ret_name) return self def squareRootOfPopulationStandardDeviation_progress(self): return self.squareRootOfPopulationStandardDeviation_progress_as("squareRootOfPopulationStandardDeviationOfProgress") def squareRootOfPopulationStandardDeviation_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.squareRootOfPopulationStandardDeviation("progress", ret_name) return self def sampleVariance_progress(self): return self.sampleVariance_progress_as("sampleVarianceOfProgress") def sampleVariance_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.sampleVariance("progress", ret_name) return self def samplePopulationVariance_progress(self): return self.samplePopulationVariance_progress_as("samplePopulationVarianceOfProgress") def samplePopulationVariance_progress_as(self, ret_name: str): - self.query.("progress", ret_name) + self.query.samplePopulationVariance("progress", ret_name) return self def group_by_id(self): self.query.group_by("id") return self def group_by_id_as(self, ret_name: str): - self.query.group_by("id") + self.query.group_by("id") return self - def group_by_name(self): self.query.group_by("name") return self def group_by_name_as(self, ret_name: str): - self.query.group_by("name") + self.query.group_by("name") return self - def group_by_code(self): self.query.group_by("code") return self def group_by_code_as(self, ret_name: str): - self.query.group_by("code") + self.query.group_by("code") return self - def group_by_color(self): self.query.group_by("color") return self def group_by_color_as(self, ret_name: str): - self.query.group_by("color") + self.query.group_by("color") return self - def group_by_display_order(self): self.query.group_by("display_order") return self def group_by_display_order_as(self, ret_name: str): - self.query.group_by("display_order") + self.query.group_by("display_order") return self - def group_by_progress(self): self.query.group_by("progress") return self def group_by_progress_as(self, ret_name: str): - self.query.group_by("progress") + self.query.group_by("progress") + return self + def group_by_platform(self): + self.query.group_by("platform") return self - + def group_by_platform_as(self, ret_name: str): + self.query.group_by("platform") + return self def group_by_version(self): self.query.group_by("version") return self def group_by_version_as(self, ret_name: str): - self.query.group_by("version") + self.query.group_by("version") + return self + def select_task_list(self): + from requests.task_request import TaskRequest + return self.select_task_list_with(TaskRequest()) + + def select_task_list_with(self, child_request): + self.query.relation_query("task_list", child_request.query) return self + def have_tasks(self): + from requests.task_request import TaskRequest + return self.with_task_list_matching(TaskRequest()) + def have_no_tasks(self): + from requests.task_request import TaskRequest + return self.without_task_list_matching(TaskRequest()) + + def with_task_list_matching(self, child_request): + child_request.query.projection = ["status"] + self.query.and_filter(in_subquery(column("id"), "Task", child_request.query)) + return self + + def without_task_list_matching(self, child_request): + child_request.query.projection = ["status"] + self.query.and_filter(not_in_subquery(column("id"), "Task", child_request.query)) + return self + def count_tasks(self): + return self.count_tasks_as("count_tasks") - async def execute_for_list(self, context, service): - req = QueryRequest(self.query) - res = await service.query(context, req) + def count_tasks_as(self, alias: str): + from requests.task_request import TaskRequest + return self.count_tasks_with(alias, TaskRequest()) - result = {"data": res.rows} - return result \ No newline at end of file + def count_tasks_with(self, alias: str, child_request): + child_request.query.count_field("id", alias) + self.query.relation_aggregates.append( + RelationAggregate("task_list", alias, child_request.query, True) + ) + return self + + + def facet_by_platform_as(self, name: str, request: QuerySelection, + include_all_facets: bool = True): + self.query.facet_by(name, "platform", request.query, include_all_facets) + return self + + +class ExecutableTaskStatusRequest: + def __init__(self, request): + self._request = request + + def comment(self, c: str): + self._request.comment(c) + return self + + def new_entity(self, context) -> TaskStatus: + 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("TaskStatus", TaskStatus()) + if not isinstance(entity, TaskStatus): + raise TypeError("entity initializer returned an incompatible TaskStatus") + return entity + + async def execute_for_result(self, context): + self = self._request + 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(context.prepare_query(self.query), _comment=self._comment, _purpose=self._purpose) + 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[TaskStatus]: + result = await self.execute_for_result(context) + query_root = EntityRoot() + return SmartList( + (TaskStatus(_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[TaskStatus]: + 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, _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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), _comment=request._comment, _purpose=request._purpose)) + 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, _comment=request._comment, _purpose=request._purpose)) + query_root = EntityRoot() + data = SmartList(TaskStatus(_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) + 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(context.prepare_query(request.query), _comment=request._comment, _purpose=request._purpose), chunk_size): + for row in chunk.rows: + yield TaskStatus(_entity_root=query_root, **row) diff --git a/examples/task_board/generated/runtime_module.py b/examples/task_board/generated/runtime_module.py new file mode 100644 index 0000000..b41620e --- /dev/null +++ b/examples/task_board/generated/runtime_module.py @@ -0,0 +1,368 @@ +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.platform import Platform +from models.task_status import TaskStatus +from models.task import Task +from models.task_execution_log import TaskExecutionLog + +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 _PlatformChecker: + 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 ("founded" not in record or _teaql_is_null(record["founded"])): + record["founded"] = Value.from_any(now) + context.record_fix_evidence("Platform", "founded", "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 "founded" not in record) or ("founded" in record and _teaql_is_null(record["founded"])): + results.append(CheckResult("required", ObjectLocation().property("founded"))) + + if (operation == "insert" and "user_email" not in record) or ("user_email" in record and _teaql_is_null(record["user_email"])): + results.append(CheckResult("required", ObjectLocation().property("user_email"))) + if "user_email" in record and _teaql_raw(record["user_email"]) is not None and len(_teaql_raw(record["user_email"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("user_email"), _teaql_raw(record["user_email"]), 100)) + + + +class _TaskStatusChecker: + 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 (operation == "insert" and "color" not in record) or ("color" in record and _teaql_is_null(record["color"])): + results.append(CheckResult("required", ObjectLocation().property("color"))) + 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 "display_order" not in record) or ("display_order" in record and _teaql_is_null(record["display_order"])): + results.append(CheckResult("required", ObjectLocation().property("display_order"))) + + if (operation == "insert" and "progress" not in record) or ("progress" in record and _teaql_is_null(record["progress"])): + results.append(CheckResult("required", ObjectLocation().property("progress"))) + + if (operation == "insert" and "platform" not in record) or ("platform" in record and _teaql_is_null(record["platform"])): + results.append(CheckResult("required", ObjectLocation().property("platform"))) + + + +class _TaskChecker: + 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 not len(_teaql_raw(record["name"])) >= 1: + results.append(CheckResult("min_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 1)) + if "name" in record and _teaql_raw(record["name"]) is not None and len(_teaql_raw(record["name"])) > 200: + results.append(CheckResult("max_length", ObjectLocation().property("name"), _teaql_raw(record["name"]), 200)) + + 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 "platform" not in record) or ("platform" in record and _teaql_is_null(record["platform"])): + results.append(CheckResult("required", ObjectLocation().property("platform"))) + + + +class _TaskExecutionLogChecker: + 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 "task" not in record) or ("task" in record and _teaql_is_null(record["task"])): + results.append(CheckResult("required", ObjectLocation().property("task"))) + + if (operation == "insert" and "action" not in record) or ("action" in record and _teaql_is_null(record["action"])): + results.append(CheckResult("required", ObjectLocation().property("action"))) + if "action" in record and _teaql_raw(record["action"]) is not None and len(_teaql_raw(record["action"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("action"), _teaql_raw(record["action"]), 100)) + + if (operation == "insert" and "detail" not in record) or ("detail" in record and _teaql_is_null(record["detail"])): + results.append(CheckResult("required", ObjectLocation().property("detail"))) + if "detail" in record and _teaql_raw(record["detail"]) is not None and len(_teaql_raw(record["detail"])) > 100: + results.append(CheckResult("max_length", ObjectLocation().property("detail"), _teaql_raw(record["detail"]), 100)) + + + +_Platform_DESCRIPTOR = (EntityDescriptor("Platform") + .audit_mask_fields([]) + .table_name("platform_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("founded", DataType.Timestamp).column_name("founded").log_policy("plain").required()).property(PropertyDescriptor("user_email", DataType.Text).column_name("user_email").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("task_status_list", "TaskStatus").local("id").foreign("platform").many()).relation(RelationDescriptor("task_list", "Task").local("id").foreign("platform").many()) +) + +_TaskStatus_DESCRIPTOR = (EntityDescriptor("TaskStatus") + .audit_mask_fields([]) + .table_name("task_status_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("code", DataType.Text).column_name("code").log_policy("plain").required()).property(PropertyDescriptor("color", DataType.Text).column_name("color").log_policy("plain").required()).property(PropertyDescriptor("display_order", DataType.Decimal).column_name("display_order").log_policy("plain").required()).property(PropertyDescriptor("progress", DataType.Decimal).column_name("progress").log_policy("plain").required()).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")).relation(RelationDescriptor("task_list", "Task").local("id").foreign("status").many()) +) + +_Task_DESCRIPTOR = (EntityDescriptor("Task") + .audit_mask_fields([]) + .table_name("task_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("name", DataType.Text).column_name("name").log_policy("plain").required()).property(PropertyDescriptor("status", DataType.I64).column_name("status").log_policy("plain").required()).property(PropertyDescriptor("platform", DataType.I64).column_name("platform").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("status", "TaskStatus").local("status").foreign("id")).relation(RelationDescriptor("platform", "Platform").local("platform").foreign("id")).relation(RelationDescriptor("task_execution_log_list", "TaskExecutionLog").local("id").foreign("task").many()) +) + +_TaskExecutionLog_DESCRIPTOR = (EntityDescriptor("TaskExecutionLog") + .audit_mask_fields(["detail"]) + .table_name("task_execution_log_data").property(PropertyDescriptor("id", DataType.I64).column_name("id").log_policy("plain").is_id().required()).property(PropertyDescriptor("task", DataType.I64).column_name("task").log_policy("plain").required()).property(PropertyDescriptor("action", DataType.Text).column_name("action").log_policy("plain").required()).property(PropertyDescriptor("detail", DataType.Text).column_name("detail").log_policy("plain").required()).property(PropertyDescriptor("version", DataType.I64).column_name("version").log_policy("plain").is_version().required()).relation(RelationDescriptor("task", "Task").local("task").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("Robot System") + platform_1.update_user_email("string()") + 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)) + task_status_1001 = await (Q.task_statuses().with_id_is(1001).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if task_status_1001 is None: + task_status_1001 = TaskStatus._teaql_new_with_fixed_id(1001) + task_status_1001.update_name("Planned") + task_status_1001.update_code("PLANNED") + task_status_1001.update_color("#94A3B8") + task_status_1001.update_display_order(10) + task_status_1001.update_progress(0) + task_status_1001.update_platform(Platform.refer(1)) + try: + await task_status_1001.audit_as('create model constant TaskStatus(1001)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + task_status_1001 = await (Q.task_statuses().with_id_is(1001).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if task_status_1001 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if task_status_1001 is None: + raise _teaql_create_error + _teaql_changed = False + if task_status_1001.name != "Planned": + task_status_1001.update_name("Planned") + _teaql_changed = True + if task_status_1001.code != "PLANNED": + task_status_1001.update_code("PLANNED") + _teaql_changed = True + if task_status_1001.color != "#94A3B8": + task_status_1001.update_color("#94A3B8") + _teaql_changed = True + if task_status_1001.displayOrder != 10: + task_status_1001.update_display_order(10) + _teaql_changed = True + if task_status_1001.progress != 0: + task_status_1001.update_progress(0) + _teaql_changed = True + if task_status_1001.platform != 1: + task_status_1001.update_platform(Platform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await task_status_1001.audit_as('reconcile model constant TaskStatus(1001)').save(context) + task_status_1002 = await (Q.task_statuses().with_id_is(1002).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if task_status_1002 is None: + task_status_1002 = TaskStatus._teaql_new_with_fixed_id(1002) + task_status_1002.update_name("Ready") + task_status_1002.update_code("READY") + task_status_1002.update_color("#3B82F6") + task_status_1002.update_display_order(20) + task_status_1002.update_progress(25) + task_status_1002.update_platform(Platform.refer(1)) + try: + await task_status_1002.audit_as('create model constant TaskStatus(1002)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + task_status_1002 = await (Q.task_statuses().with_id_is(1002).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if task_status_1002 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if task_status_1002 is None: + raise _teaql_create_error + _teaql_changed = False + if task_status_1002.name != "Ready": + task_status_1002.update_name("Ready") + _teaql_changed = True + if task_status_1002.code != "READY": + task_status_1002.update_code("READY") + _teaql_changed = True + if task_status_1002.color != "#3B82F6": + task_status_1002.update_color("#3B82F6") + _teaql_changed = True + if task_status_1002.displayOrder != 20: + task_status_1002.update_display_order(20) + _teaql_changed = True + if task_status_1002.progress != 25: + task_status_1002.update_progress(25) + _teaql_changed = True + if task_status_1002.platform != 1: + task_status_1002.update_platform(Platform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await task_status_1002.audit_as('reconcile model constant TaskStatus(1002)').save(context) + task_status_1003 = await (Q.task_statuses().with_id_is(1003).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if task_status_1003 is None: + task_status_1003 = TaskStatus._teaql_new_with_fixed_id(1003) + task_status_1003.update_name("Executing") + task_status_1003.update_code("EXECUTING") + task_status_1003.update_color("#F59E0B") + task_status_1003.update_display_order(30) + task_status_1003.update_progress(50) + task_status_1003.update_platform(Platform.refer(1)) + try: + await task_status_1003.audit_as('create model constant TaskStatus(1003)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + task_status_1003 = await (Q.task_statuses().with_id_is(1003).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if task_status_1003 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if task_status_1003 is None: + raise _teaql_create_error + _teaql_changed = False + if task_status_1003.name != "Executing": + task_status_1003.update_name("Executing") + _teaql_changed = True + if task_status_1003.code != "EXECUTING": + task_status_1003.update_code("EXECUTING") + _teaql_changed = True + if task_status_1003.color != "#F59E0B": + task_status_1003.update_color("#F59E0B") + _teaql_changed = True + if task_status_1003.displayOrder != 30: + task_status_1003.update_display_order(30) + _teaql_changed = True + if task_status_1003.progress != 50: + task_status_1003.update_progress(50) + _teaql_changed = True + if task_status_1003.platform != 1: + task_status_1003.update_platform(Platform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await task_status_1003.audit_as('reconcile model constant TaskStatus(1003)').save(context) + task_status_1004 = await (Q.task_statuses().with_id_is(1004).comment('what: locate generated bootstrap entity').purpose('why: idempotent runtime bootstrap').execute_for_one(context)) + if task_status_1004 is None: + task_status_1004 = TaskStatus._teaql_new_with_fixed_id(1004) + task_status_1004.update_name("Verified") + task_status_1004.update_code("VERIFIED") + task_status_1004.update_color("#16A34A") + task_status_1004.update_display_order(40) + task_status_1004.update_progress(100) + task_status_1004.update_platform(Platform.refer(1)) + try: + await task_status_1004.audit_as('create model constant TaskStatus(1004)').save(context) + except Exception as _teaql_create_error: + for _teaql_attempt in range(5): + task_status_1004 = await (Q.task_statuses().with_id_is(1004).comment('what: recover concurrent bootstrap').purpose('why: make generated bootstrap idempotent').execute_for_one(context)) + if task_status_1004 is not None: + break + if _teaql_attempt < 4: + await asyncio.sleep((_teaql_attempt + 1) * 0.01) + if task_status_1004 is None: + raise _teaql_create_error + _teaql_changed = False + if task_status_1004.name != "Verified": + task_status_1004.update_name("Verified") + _teaql_changed = True + if task_status_1004.code != "VERIFIED": + task_status_1004.update_code("VERIFIED") + _teaql_changed = True + if task_status_1004.color != "#16A34A": + task_status_1004.update_color("#16A34A") + _teaql_changed = True + if task_status_1004.displayOrder != 40: + task_status_1004.update_display_order(40) + _teaql_changed = True + if task_status_1004.progress != 100: + task_status_1004.update_progress(100) + _teaql_changed = True + if task_status_1004.platform != 1: + task_status_1004.update_platform(Platform.refer(1)) + _teaql_changed = True + if _teaql_changed: + await task_status_1004.audit_as('reconcile model constant TaskStatus(1004)').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(Platform) + .schema_entity(_Platform_DESCRIPTOR) + .checker("Platform", _PlatformChecker()) + .wire_metadata("Platform", create_wire_entity_metadata("Platform", ["id", "name", "founded", "user_email", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "founded": ["founded"], "user_email": ["user_email"], "version": ["version"]})).entity(TaskStatus) + .schema_entity(_TaskStatus_DESCRIPTOR) + .checker("TaskStatus", _TaskStatusChecker()) + .wire_metadata("TaskStatus", create_wire_entity_metadata("TaskStatus", ["id", "name", "code", "color", "display_order", "progress", "platform", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "code": ["code"], "color": ["color"], "display_order": ["display_order"], "progress": ["progress"], "platform": ["platform"], "version": ["version"]})).entity(Task) + .schema_entity(_Task_DESCRIPTOR) + .checker("Task", _TaskChecker()) + .wire_metadata("Task", create_wire_entity_metadata("Task", ["id", "name", "status", "platform", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "name": ["name"], "status": ["status"], "platform": ["platform"], "version": ["version"]})).entity(TaskExecutionLog) + .schema_entity(_TaskExecutionLog_DESCRIPTOR) + .checker("TaskExecutionLog", _TaskExecutionLogChecker()) + .wire_metadata("TaskExecutionLog", create_wire_entity_metadata("TaskExecutionLog", ["id", "task", "action", "detail", "version"], JsonFieldNamingProfile.CAMEL_CASE, {"id": ["id"], "task": ["task"], "action": ["action"], "detail": ["detail"], "version": ["version"]})) + .generated_bootstrap(_ensure_generated_bootstrap) +) \ No newline at end of file diff --git a/examples/task_board/generated/teaql-i18n.json b/examples/task_board/generated/teaql-i18n.json new file mode 100644 index 0000000..993843e --- /dev/null +++ b/examples/task_board/generated/teaql-i18n.json @@ -0,0 +1,123 @@ +{ + "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.platform.name": "Name", + "entity.task": "Task", + "property.taskStatus.progress": "Progress", + "entity.taskStatus": "Task Status", + "property.taskStatus.name": "Name", + "property.task.name": "Name", + "property.taskExecutionLog.action": "Action", + "entity.taskExecutionLog": "Task Execution Log", + "property.platform.founded": "Founded", + "property.taskStatus.color": "Color", + "property.taskExecutionLog.id": "Id", + "property.taskExecutionLog.version": "Version", + "property.taskStatus.id": "Id", + "entity.platform": "Platform", + "property.platform.id": "Id", + "property.taskStatus.code": "Code", + "property.task.id": "Id", + "property.taskStatus.displayOrder": "Display Order", + "property.taskExecutionLog.task": "Task", + "property.taskExecutionLog.detail": "Detail", + "property.platform.userEmail": "User Email", + "property.task.platform": "Platform", + "property.task.status": "Status", + "property.taskStatus.version": "Version", + "property.task.version": "Version", + "property.taskStatus.platform": "Platform", + "property.platform.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/task_board/main.py b/examples/task_board/main.py index c0bc752..7043005 100644 --- a/examples/task_board/main.py +++ b/examples/task_board/main.py @@ -1,159 +1,64 @@ +"""Current generated library + local runtime; no manual DDL or bootstrap writes.""" import asyncio -import aiosqlite import os -import time -from teaql.provider.sqlite import create_sqlite_service, SimpleSchemaProvider -from teaql.core.meta import EntityDescriptor, PropertyDescriptor -from teaql.core.value import DataType, Value, Timestamp -from teaql.data_service import QueryRequest -from teaql.runtime.context import UserContext +from pathlib import Path +import sys +import tempfile -from generated.models.platform import Platform -from generated.models.task_status import TaskStatus -from generated.models.task import Task -from generated.models.task_execution_log import TaskExecutionLog +sys.path.insert(0, str(Path(__file__).resolve().with_name("generated"))) -from generated.requests.task_request import TaskRequest -from generated.requests.task_execution_log_request import TaskExecutionLogRequest +from E import E +from Q import Q +from models.task import Task +from models.task_execution_log import TaskExecutionLog +from runtime_module import GENERATED_RUNTIME_MODULE +from teaql.data_service import SQLiteTeaQLClient +from teaql.runtime import UserContext async def main(): - print("Setting up Schema Provider...") - provider = SimpleSchemaProvider() + database = os.environ.get("TEAQL_TASK_BOARD_DB") + if not database: + with tempfile.NamedTemporaryFile(prefix="teaql-task-board-", suffix=".db", delete=False) as file: + database = file.name + client = SQLiteTeaQLClient(database) + context = UserContext.new().install(GENERATED_RUNTIME_MODULE).insert_resource("dataService", client) + try: + await context.ensure_schema() + await context.ensure_schema() + task = Task(name="Build Robot Arm", platform=1).update_status_to_planned() + await task.audit_as("create demo robot task").save(context) + loaded = await (Q.tasks().with_id_is(task.id).limit(1) + .select_status_with(Q.task_statuses().limit(1)) + .comment("load task and planned status").purpose("review complete task before editing") + .execute_for_one(context)) + assert E.task(loaded).name().eval() == "Build Robot Arm" + assert E.task(loaded).status().name().eval() == "Planned" + await loaded.update_name("Build Robot Arm V2").audit_as("rename demo robot task").save(context) + updated = await (Q.tasks().with_id_is(task.id).limit(1) + .comment("read renamed task").purpose("verify persisted mutation") + .execute_for_one(context)) + assert updated.name == "Build Robot Arm V2" + page = await (Q.tasks().with_id_is(task.id) + .comment("page renamed task").purpose("verify paginated task intent") + .execute_for_page(context, 0, 1)) + assert page.total_count == 1 and page.data[0].name == "Build Robot Arm V2" + streamed = [] + async for row in (Q.tasks().with_id_is(task.id) + .comment("stream renamed task").purpose("verify streamed task intent") + .execute_for_stream(context, chunk_size=1)): + streamed.append(row) + assert len(streamed) == 1 and streamed[0].name == "Build Robot Arm V2" + entry = TaskExecutionLog(task=task.id, action="RENAME", detail="PRIVATE-TASK-DETAIL") + await entry.audit_as("record PRIVATE-TASK-DETAIL rename evidence").save(context) + rows = await (Q.task_execution_logs().with_id_is(entry.id).with_detail_is("PRIVATE-TASK-DETAIL").limit(1) + .comment("read PRIVATE-TASK-DETAIL evidence").purpose("verify PRIVATE-TASK-DETAIL is persisted") + .execute_for_list(context)) + assert len(rows) == 1 and E.task_execution_log(rows[0]).detail().eval() == "PRIVATE-TASK-DETAIL" + print("PASS task board: governed Q/E/mutation and masked detail") + finally: + await client.close() - # 1. Platform - platform_entity = EntityDescriptor("Platform")\ - .table_name("platform")\ - .property(PropertyDescriptor("id", DataType.I64).is_id())\ - .property(PropertyDescriptor("name", DataType.Text))\ - .property(PropertyDescriptor("founded", DataType.Timestamp))\ - .property(PropertyDescriptor("user_email", DataType.Text))\ - .property(PropertyDescriptor("version", DataType.I64).is_version()) - provider.register_entity(platform_entity) - - # 2. TaskStatus - task_status_entity = EntityDescriptor("TaskStatus")\ - .table_name("task_status")\ - .property(PropertyDescriptor("id", DataType.I64).is_id())\ - .property(PropertyDescriptor("name", DataType.Text))\ - .property(PropertyDescriptor("code", DataType.Text))\ - .property(PropertyDescriptor("color", DataType.Text))\ - .property(PropertyDescriptor("displayOrder", DataType.I64))\ - .property(PropertyDescriptor("progress", DataType.I64))\ - .property(PropertyDescriptor("platform", DataType.I64))\ - .property(PropertyDescriptor("version", DataType.I64).is_version()) - provider.register_entity(task_status_entity) - - # 3. Task - task_entity = EntityDescriptor("Task")\ - .table_name("task")\ - .property(PropertyDescriptor("id", DataType.I64).is_id())\ - .property(PropertyDescriptor("name", DataType.Text))\ - .property(PropertyDescriptor("status", DataType.I64))\ - .property(PropertyDescriptor("platform", DataType.I64))\ - .property(PropertyDescriptor("version", DataType.I64).is_version()) - provider.register_entity(task_entity) - - # 4. TaskExecutionLog - log_entity = EntityDescriptor("TaskExecutionLog")\ - .table_name("task_execution_log")\ - .property(PropertyDescriptor("id", DataType.I64).is_id())\ - .property(PropertyDescriptor("task", DataType.I64))\ - .property(PropertyDescriptor("action", DataType.Text))\ - .property(PropertyDescriptor("detail", DataType.Text))\ - .property(PropertyDescriptor("version", DataType.I64).is_version()) - provider.register_entity(log_entity) - - db_path = os.environ.get("TEAQL_TASK_BOARD_DB", "task_board.db") - print(f"Creating sqlite service at {db_path}...") - service = create_sqlite_service(db_path, provider) - - print("Initializing Database Schema...") - context = UserContext.new() - for e in provider.entities.values(): - context.register_entity(e) - - context.with_schema_provider(service) - await context.ensure_schema() - - print("Performing CRUD operations...") - - # CREATE Platform - print("Creating Platform...") - platform = Platform( - id=1, - name="Main Platform", - founded=Timestamp(int(time.time() * 1000)), - userEmail="admin@robot.com", - version=1 - ) - platform._action = "Create" - await platform.save(context, service) - - # CREATE TaskStatus - print("Creating TaskStatus...") - status = TaskStatus( - id=1, - name="Planned", - code="PLANNED", - color="#94A3B8", - displayOrder=10, - progress=0, - platform=1, - version=1 - ) - status._action = "Create" - await status.save(context, service) - - # CREATE Task - print("Creating Task...") - task = Task( - id=1, - name="Build Robot Arm", - status=1, - platform=1, - version=1 - ) - task._action = "Create" - await task.save(context, service) - - # READ Task - print("Reading Tasks...") - task_req = TaskRequest() - res = await task_req.execute_for_list(context, service) - for row in res["data"]: - print(" - Task:", row) - - # UPDATE Task - print("Updating Task...") - task.name = "Build Robot Leg" - task._action = "Update" - await task.save(context, service) - - # Verify Update - res = await task_req.execute_for_list(context, service) - print(" - Updated Task:", res["data"][0]) - - # INSERT TaskExecutionLog - print("Inserting TaskExecutionLog...") - log = TaskExecutionLog( - id=1, - task=1, - action="Updated Task", - detail="Changed arm to leg", - version=1 - ) - log._action = "Create" - await log.save(context, service) - - # QUERY Log - print("Reading Logs...") - log_req = TaskExecutionLogRequest() - res = await log_req.execute_for_list(context, service) - for row in res["data"]: - print(" - Log:", row) - - print("Success!") if __name__ == "__main__": asyncio.run(main()) diff --git a/examples/task_board/test_task_board.py b/examples/task_board/test_task_board.py new file mode 100644 index 0000000..4b1dec4 --- /dev/null +++ b/examples/task_board/test_task_board.py @@ -0,0 +1,47 @@ +"""Run the real example; do not infer governance from process exit zero alone.""" +import os +from pathlib import Path +import subprocess +import sys +import tempfile +import unittest + + +class TaskBoardLogContractTest(unittest.TestCase): + def test_default_logs_have_intent_and_mask_business_values(self): + repo = Path(__file__).resolve().parents[2] + env = dict(os.environ) + env.pop("TEAQL_ALLOW_SENSITIVE_PLAINTEXT_LOGS", None) + env["PYTHONPATH"] = str(repo / "src") + with tempfile.TemporaryDirectory(prefix="teaql-task-board-test-") as temporary: + env["TEAQL_TASK_BOARD_DB"] = str(Path(temporary) / "task-board.db") + result = subprocess.run([sys.executable, str(Path(__file__).with_name("main.py"))], + env=env, capture_output=True, text=True, timeout=90) + output = result.stdout + result.stderr + self.assertEqual(0, result.returncode, output) + queries = writes = 0 + for line in output.splitlines(): + if not line.startswith("[TeaQL SQL]"): + continue + if "[select]" in line or "[query]" in line: + queries += 1 + self.assertNotIn("comment=None", line) + self.assertNotIn("purpose=None", line) + self.assertNotIn("comment= purpose=", line) + self.assertNotIn("purpose= auditReason=", line) + elif any(f"[{kind}]" in line for kind in ("insert", "update", "delete")): + writes += 1 + self.assertNotIn("auditReason=None", line) + self.assertNotIn("auditReason= tracePath=", line) + self.assertGreater(queries, 0, output) + self.assertGreater(writes, 0, output) + self.assertIn("PASS task board: governed Q/E/mutation and masked detail", output) + self.assertNotIn("PRIVATE-TASK-DETAIL", output) + self.assertIn("'PR***************IL' /* masked */", output) + self.assertIn("'RENAME'", output) + self.assertIn("comment='page renamed task' purpose='verify paginated task intent'", output) + self.assertIn("comment='stream renamed task' purpose='verify streamed task intent'", output) + + +if __name__ == "__main__": + unittest.main() diff --git a/scripts/verify-examples.sh b/scripts/verify-examples.sh index c87a3b4..0dc7274 100755 --- a/scripts/verify-examples.sh +++ b/scripts/verify-examples.sh @@ -21,12 +21,16 @@ if rg -l 'include\s*=.*teaql\*' "$repo/examples" --glob pyproject.toml >/dev/nul fi PYTHONPATH="$repo/examples/conformance:$repo/src" python -m app.main +PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/conformance" -p 'test_sql_log_intent.py' -v PYTHONPATH="$repo/examples/school-management:$repo/src" python -m app.main +PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/school-management" -p 'test_sql_log_intent.py' -v order_management_tmp="$(mktemp -d)" task_board_tmp="$(mktemp -d)" 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" +PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/order-management" -p 'test_sql_log_intent.py' -v TEAQL_TASK_BOARD_DB="$task_board_tmp/task_board.db" PYTHONPATH="$repo/examples/task_board:$repo/src" python "$repo/examples/task_board/main.py" +PYTHONPATH="$repo/src" python -m unittest discover -s "$repo/examples/task_board" -p 'test_task_board.py' -v echo "PASS: all Python examples" diff --git a/src/teaql/core/meta.py b/src/teaql/core/meta.py index 630d570..18cedca 100644 --- a/src/teaql/core/meta.py +++ b/src/teaql/core/meta.py @@ -6,6 +6,13 @@ def __init__(self, name: str, property_type: str = "String"): self._is_id = False self._is_version = False self.nullable = True + self.log_policy_val = 'unknown' + def log_policy(self, policy): + """Trusted schema policy; incoming requests cannot override it.""" + if policy not in ('plain', 'masked', 'credential', 'unknown'): + raise ValueError('invalid SQL parameter log policy') + self.log_policy_val = policy + return self def column_name(self, name): self.column_name_val = name return self diff --git a/src/teaql/data_service/__init__.py b/src/teaql/data_service/__init__.py index 3a032c3..5bcbe59 100644 --- a/src/teaql/data_service/__init__.py +++ b/src/teaql/data_service/__init__.py @@ -71,6 +71,11 @@ class ExecutionMetadata: audit_reason: Optional[str] = None backend_request_id: Optional[str] = None debug_query: Optional[str] = None + database_kind: Any = None + parameter_log_policies: Optional[List[str]] = None + sql_origin: Optional[str] = None + # Statement/cursor termination only, not transaction commit. + execution_outcome: Optional[str] = None @dataclass diff --git a/src/teaql/provider/postgres/dialect.py b/src/teaql/provider/postgres/dialect.py index 292c67b..1a93e0f 100644 --- a/src/teaql/provider/postgres/dialect.py +++ b/src/teaql/provider/postgres/dialect.py @@ -5,7 +5,7 @@ class PostgresDialect(SqlDialect): def kind(self) -> DatabaseKind: - return DatabaseKind.Postgres + return DatabaseKind.PostgreSql def quote_ident(self, ident: str) -> str: return quote_identifier_if_needed(ident, '"') diff --git a/src/teaql/runtime/audit.py b/src/teaql/runtime/audit.py index a1ff23a..c52cbd5 100644 --- a/src/teaql/runtime/audit.py +++ b/src/teaql/runtime/audit.py @@ -35,15 +35,15 @@ def safe(self, mask_fields: List[str], max_length: Optional[int]) -> "SafeAuditE fields = [] for change in self.changes: value = None if change.new_value is None else str(getattr(change.new_value, "val", change.new_value)) - masked = (credential_name(change.field) - or payload_has_credentials(change.new_value) - or payload_has_credentials(change.old_value) - or (change.field in mask_fields and not allow)) + credential = (credential_name(change.field) + or payload_has_credentials(change.new_value) + or payload_has_credentials(change.old_value)) + masked = credential or (change.field in mask_fields and not allow) if masked: secrets.extend(value_strings(change.old_value)) secrets.extend(value_strings(change.new_value)) if value is not None and masked: - value = REDACTED + value = REDACTED if credential else _mask(value) truncated = value is not None and max_length is not None and len(value) > max_length if truncated: value = "*" * max_length if max_length <= 3 else value[:max_length - 3] + "..." @@ -74,7 +74,8 @@ class SafeAuditEvent: def _mask(value: str) -> str: - if len(value) < 8: + # Unicode scalar length, ASCII digits: the same contract as Rust and Go. + if len(value) < 8 or (value.isascii() and value.isdigit()): return "*" * len(value) return value[:2] + "*" * (len(value) - 4) + value[-2:] diff --git a/src/teaql/runtime/context.py b/src/teaql/runtime/context.py index ff84690..fa928fd 100644 --- a/src/teaql/runtime/context.py +++ b/src/teaql/runtime/context.py @@ -705,6 +705,10 @@ def language(self) -> Any: return self.get_resource("locale") or Locale.ENGLISH def record_metadata_log(self, metadata: Any): + self._record_metadata_log(metadata) + + def _record_metadata_log(self, metadata: Any, *, intent_source=None): + """Internal statement plumbing: source bindings never reach sinks/buffers.""" op = SqlLogOperation.Select op_str = str(getattr(metadata, 'operation', '')).lower() if 'insert' in op_str: op = SqlLogOperation.Insert @@ -726,6 +730,10 @@ def record_metadata_log(self, metadata: Any): params=list(getattr(metadata, 'parameters', [])), debug_sql=getattr(metadata, 'debug_query', '') or '', pretty_sql=getattr(metadata, 'debug_query', '') or '', + database_kind=getattr(metadata, 'database_kind', None), + parameter_log_policies=getattr(metadata, 'parameter_log_policies', None), + sql_origin=getattr(metadata, 'sql_origin', None), + execution_outcome=getattr(metadata, 'execution_outcome', None), started_at=started_at, ended_at=ended_at, elapsed=ended_at - started_at, @@ -740,7 +748,7 @@ def record_metadata_log(self, metadata: Any): entry.result_summary = f"{entry.affected_rows} rows affected" from .log_privacy import sql_log_projection - entry = sql_log_projection(entry) + entry = sql_log_projection(entry, _intent_source=intent_source) logs = self.sql_logs() logs.append(entry) self._resources["sql_logs"] = logs @@ -760,7 +768,8 @@ def record_sql_log(self, operation: Any, query: Any, started_at: Any, ended_at: if not self.sql_log_options().enabled_for(operation): return - debug_sql = getattr(query, 'debug_sql', lambda *args: "")() if hasattr(query, 'debug_sql') else getattr(query, 'sql', "") + # Projection renders only safe values. Never materialize plaintext first. + debug_sql = '' entry = SqlLogEntry( operation=operation, @@ -772,6 +781,9 @@ def record_sql_log(self, operation: Any, query: Any, started_at: Any, ended_at: params=getattr(query, 'params', []), debug_sql=debug_sql, pretty_sql=debug_sql, + database_kind=getattr(query, 'database_kind', None), + parameter_log_policies=getattr(query, 'parameter_log_policies', None), + sql_origin=getattr(query, 'sql_origin', None), started_at=started_at, ended_at=ended_at, elapsed=elapsed, @@ -1035,6 +1047,13 @@ class SqlLogEntry: result_type: Optional[str] affected_rows: Optional[int] result_summary: str + database_kind: Any = None + parameter_log_policies: Optional[List[str]] = None + sql_origin: Optional[str] = None + masked_parameters: Optional[List[bool]] = None + log_mode: Optional[str] = None + omission_reason: Optional[str] = None + execution_outcome: Optional[str] = None class DiagnosticSqlLogSink: """Policy-projected SQL destination; the text sink is installed by default.""" @@ -1051,9 +1070,9 @@ def write(self, entry: SqlLogEntry) -> None: elapsed_us = int(entry.elapsed.total_seconds() * 1_000_000) if entry.elapsed else 0 self._writer( f"[TeaQL SQL][{entry.operation.name.lower()}][{elapsed_us}us] " - f"{entry.result_summary} comment={entry.comment!r} purpose={entry.purpose!r} " + f"{entry.result_summary} outcome={entry.execution_outcome or 'unknown'} comment={entry.comment!r} purpose={entry.purpose!r} " f"auditReason={entry.audit_reason!r} tracePath={entry.trace_path!r}\n" - f"Parameterized SQL: {entry.sql} params={entry.params!r}\n" + f"SQL omission reason: {entry.omission_reason or 'none'}\n" f"Debug SQL: {entry.debug_sql}" ) diff --git a/src/teaql/runtime/log_privacy.py b/src/teaql/runtime/log_privacy.py index 322fd3b..1f71527 100644 --- a/src/teaql/runtime/log_privacy.py +++ b/src/teaql/runtime/log_privacy.py @@ -3,6 +3,8 @@ import logging import os import re +import weakref +from copy import deepcopy from dataclasses import fields, is_dataclass, replace from functools import lru_cache @@ -10,6 +12,11 @@ PLAINTEXT_ACK = "I_UNDERSTAND_SENSITIVE_DATA_MAY_BE_WRITTEN_TO_DISK" REDACTED = "[REDACTED]" SQL_REDACTED = "[REDACTED SQL; NOT REPLAYABLE]" +DEBUG_LABEL = "-- TeaQL DEBUG PLAINTEXT; EXPLICIT OPT-IN\n" + + +def label_debug_sql(sql): + return sql if not sql or sql.startswith(DEBUG_LABEL) else DEBUG_LABEL + sql @lru_cache(maxsize=1) @@ -54,40 +61,134 @@ def value_strings(value): return [] if value is None or value == "" else [str(value)] -def scrub(value, secrets): +def scrub(value, secrets, hide_all=False): """Copy annotations, removing known values even when embedded in intent.""" if isinstance(value, str): + if hide_all: + return REDACTED if value else value for secret in sorted(set(secrets), key=len, reverse=True): value = value.replace(secret, REDACTED) return value if isinstance(value, list): - return [scrub(v, secrets) for v in value] + return [scrub(v, secrets, hide_all) for v in value] if isinstance(value, tuple): - return tuple(scrub(v, secrets) for v in value) + return tuple(scrub(v, secrets, hide_all) for v in value) if isinstance(value, dict): - return {k: scrub(v, secrets) for k, v in value.items()} + return {k: scrub(v, secrets, hide_all) for k, v in value.items()} if is_dataclass(value): - return replace(value, **{f.name: scrub(getattr(value, f.name), secrets) + return replace(value, **{f.name: scrub(getattr(value, f.name), secrets, hide_all) for f in fields(value) if f.init}) return value -def sql_log_projection(entry): - # SQL has no per-parameter field provenance. A credential-bearing query - # therefore suppresses its whole payload, including on debug opt-in. - credentials = (credential_name(entry.sql) or credential_name(entry.debug_sql) - or payload_has_credentials(entry.params)) - allow = plaintext_enabled() and not credentials - if allow: - return replace(entry, params=list(entry.params), trace_path=list(entry.trace_path)) - secrets = value_strings(entry.params) - sql = entry.sql - # Arbitrary literal SQL is not safely redactable across dialects. Preserve - # parameterized shape only when there are no literals/comments to expose. - if any(token in sql for token in ("'", '"', "`", "--", "/*", "$")) or re.search(r"\b\d+\b", sql): - sql = SQL_REDACTED - return replace(entry, sql=scrub(sql, secrets), params=[None] * len(entry.params), - debug_sql=SQL_REDACTED, pretty_sql=SQL_REDACTED, - comment=scrub(entry.comment, secrets), purpose=scrub(entry.purpose, secrets), - audit_reason=scrub(entry.audit_reason, secrets), - trace_path=scrub(entry.trace_path, secrets)) +_projections = {} + +def _remember_projection(projected, allow, alternative=None): + key = id(projected) + _projections[key] = (weakref.ref(projected, lambda _: _projections.pop(key, None)), + allow, deepcopy(projected), alternative) + return projected + +def _business_mask(value): + from .audit import _mask + value = getattr(value, 'val', value) + if value is None: + return None + if isinstance(value, (tuple, list)): + return [_business_mask(child) for child in value] + if isinstance(value, dict): + return REDACTED + return _mask(str(value)) + +def _binding_policies(source): + supplied = source.parameter_log_policies + valid = supplied is not None and len(supplied) == len(source.params) + credentials = credential_name(source.sql) and (source.sql_origin != 'generated' or not valid) + policies = [] + for index, value in enumerate(source.params): + policy = supplied[index] if valid else 'unknown' + if credentials or payload_has_credentials(value): + policy = 'credential' + policies.append(policy if policy in ('plain', 'masked', 'credential') else 'unknown') + return policies + + +def _is_masked(policy, allow): + return policy in ('credential', 'unknown') or (not allow and policy != 'plain') + + +def sql_log_projection(entry, *, _intent_source=None): + """Source bindings are call-local runtime plumbing, never stored on a log entry.""" + allow = plaintext_enabled() and entry.log_mode != 'masked' + prior = _projections.get(id(entry)) + # Log entries are mutable: reuse only an unchanged projection. Weak references + # keep this idempotence bookkeeping from retaining a long-running log history. + if prior and prior[0]() is entry and prior[2] == entry: + if not prior[1] or allow: + return entry + # Entries are mutable. Do not expose the cached safe alternative itself. + if prior[3] is not None: + return _remember_projection(deepcopy(prior[3]), False) + projected = _project_with_policy(entry, allow, _intent_source) + alternative = _project_with_policy(entry, False, _intent_source) if allow else None + return _remember_projection(projected, allow, alternative) + + +def _project_with_policy(entry, allow, intent_source): + from teaql.sql.types import DatabaseKind, render_log_sql, _sql_literal + from teaql.core.value import Value + supplied = entry.parameter_log_policies + valid = supplied is None or len(supplied) == len(entry.params) + # Compiler-owned bindings already identify individual credential fields. A + # credential column elsewhere in the projection must not hide ordinary binds. + # Unclassified/custom statements retain the conservative whole-SQL fallback. + credentials = credential_name(entry.sql) and ( + entry.sql_origin != 'generated' or supplied is None or not valid) + policies = _binding_policies(entry) + masked = [_is_masked(policy, allow) for policy in policies] + safe_values = [(_business_mask(value) if policies[index] == 'masked' else REDACTED) + if masked[index] else deepcopy(value) for index, value in enumerate(entry.params)] + secrets = [text for index, value in enumerate(entry.params) if masked[index] + for text in value_strings(value)] + if intent_source is not None: + source_policies = _binding_policies(intent_source) + secrets.extend(text for index, value in enumerate(intent_source.params) + if _is_masked(source_policies[index], allow) for text in value_strings(value)) + # A copied/changed debug record has lost its reliable private alternative. + # Its inherited intent may mention values absent from its own SQL bindings. + unknown_debug_intent = not allow and entry.log_mode == 'debug-plaintext' and intent_source is None + def intent(value): + return scrub(value, secrets, hide_all=unknown_debug_intent) + bare = re.sub(r'\$[0-9]+', '?', entry.sql) + unsafe = ((not allow or credentials) and entry.sql_origin != 'generated' and + bool(re.search(r"['\"`$]|--|/\*|\b\d+\b|:[A-Za-z_]", bare))) + reason = 'untrusted-literal-sql' if unsafe else 'policy-count-mismatch' if not valid else None + rendered = SQL_REDACTED + kind = entry.database_kind or (DatabaseKind.PostgreSql if re.search(r'\$[0-9]+', entry.sql) + else DatabaseKind.MySql if '%s' in entry.sql else DatabaseKind.Sqlite) + if reason is None: + try: + def literal(index): + value = safe_values[index] + value = value if isinstance(value, Value) else Value.from_any(value) + return _sql_literal(value, kind) + (' /* masked */' if masked[index] else '') + rendered = render_log_sql(entry.sql, safe_values, kind, literal) + prefix = (('-- TeaQL DEBUG PLAINTEXT; EXPLICIT OPT-IN; PARTIALLY MASKED; NOT REPLAYABLE\n' + if any(masked) else DEBUG_LABEL) if allow else '-- TeaQL MASKED; NOT REPLAYABLE\n') + rendered = prefix + rendered + except Exception: + # Renderer failures must not echo source SQL, values, or exception text. + rendered = SQL_REDACTED + reason = 'unsupported-or-mismatched-bindings' + projected = replace(entry, sql=SQL_REDACTED if unsafe else entry.sql + if entry.sql_origin == 'generated' else scrub(entry.sql, secrets), + params=safe_values, parameter_log_policies=policies, masked_parameters=masked, + debug_sql=rendered, pretty_sql=rendered, omission_reason=reason, + log_mode='debug-plaintext' if allow else 'masked', + comment=intent(entry.comment), purpose=intent(entry.purpose), + audit_reason=intent(entry.audit_reason), + result_summary=(f'{entry.result_count} rows returned' if entry.result_count is not None + else f'{entry.affected_rows} rows affected' if entry.affected_rows is not None + else intent(entry.result_summary)), + trace_path=intent(entry.trace_path)) + return projected diff --git a/src/teaql/sql/dialect.py b/src/teaql/sql/dialect.py index e70ae5e..a7827c1 100644 --- a/src/teaql/sql/dialect.py +++ b/src/teaql/sql/dialect.py @@ -15,7 +15,7 @@ from teaql.core.value import Value, DataType from teaql.core.meta import EntityDescriptor, PropertyDescriptor from .types import ( - DatabaseKind, CompiledQuery, SqlCompileError, + DatabaseKind, CompiledQuery, SQLBindings, SqlCompileError, UnknownEntityError, UnknownFieldError, EmptyInListError, MissingIdPropertyError, MissingVersionPropertyError, EmptyMutationError, InvalidRecoverVersionError, @@ -116,12 +116,15 @@ def compile_create_table(self, entity: EntityDescriptor) -> str: return f"CREATE TABLE IF NOT EXISTS {self.quote_ident(table_name)} ({columns_str})" def compile_select(self, entity: EntityDescriptor, query: SelectQuery) -> CompiledQuery: - params: List[Value] = [] + params = SQLBindings() sql = self.compile_select_sql(entity, query, params) query_comment = query.comment_text return CompiledQuery(sql=sql, params=params, comment=query_comment) def compile_select_sql(self, entity: EntityDescriptor, query: SelectQuery, params: List[Value]) -> str: + if isinstance(params, SQLBindings) and (query.raw_sql is not None or query.raw_sql_search_criteria + or query.raw_projections or query.dynamic_properties): + params.trusted = False if query.raw_sql is not None: return query.raw_sql @@ -149,7 +152,7 @@ def compile_select_sql(self, entity: EntityDescriptor, query: SelectQuery, param search_parts = [] for prop in getattr(entity, 'properties', []): if prop.property_type == DataType.Text: - params.append(Value.Text(like_value)) + self.bind_field(params, Value.Text(like_value), entity, prop.name) search_parts.append(f"{self.quote_ident(prop.column_name_val)} LIKE {self.placeholder(len(params))}") if search_parts: where_parts.append("(" + " OR ".join(search_parts) + ")") @@ -189,7 +192,7 @@ def compile_select_sql(self, entity: EntityDescriptor, query: SelectQuery, param def compile_insert(self, entity: EntityDescriptor, command: InsertCommand) -> CompiledQuery: columns = [] placeholders = [] - params = [] + params = SQLBindings() for prop in getattr(entity, 'properties', []): prop_name = getattr(prop, 'name', None) if prop_name in command.values: @@ -198,7 +201,7 @@ def compile_insert(self, entity: EntityDescriptor, command: InsertCommand) -> Co if val._data is None: ptype = getattr(prop, 'property_type', None) or getattr(prop, 'data_type', DataType.Text) val = Value.TypedNull(ptype) - params.append(val) + self.bind_field(params, val, entity, prop_name) placeholders.append(self.placeholder(len(params))) if not columns: @@ -215,7 +218,7 @@ def compile_update(self, entity: EntityDescriptor, command: UpdateCommand) -> Co raise MissingIdPropertyError(entity._name) assignments = [] - params = [] + params = SQLBindings() for prop in getattr(entity, 'properties', []): if getattr(prop, '_is_id', False) or getattr(prop, 'is_id_val', False): continue @@ -228,24 +231,24 @@ def compile_update(self, entity: EntityDescriptor, command: UpdateCommand) -> Co if val._data is None: ptype = getattr(prop, 'property_type', None) or getattr(prop, 'data_type', DataType.Text) val = Value.TypedNull(ptype) - params.append(val) + self.bind_field(params, val, entity, prop_name) assignments.append(f"{self.quote_ident(prop.column_name_val)} = {self.placeholder(len(params))}") version_property = next((p for p in getattr(entity, 'properties', []) if getattr(p, '_is_version', False) or getattr(p, 'is_version_val', False)), None) if command.expected_version_val is not None: if not version_property: raise MissingVersionPropertyError(entity._name) - params.append(Value.I64(command.expected_version_val + 1)) + self.bind_field(params, Value.I64(command.expected_version_val + 1), entity, version_property.name) assignments.append(f"{self.quote_ident(version_property.column_name_val)} = {self.placeholder(len(params))}") if not assignments: raise EmptyMutationError("update") - params.append(command.id) + self.bind_field(params, command.id, entity, id_property.name) predicates = [f"{self.quote_ident(id_property.column_name_val)} = {self.placeholder(len(params))}"] if command.expected_version_val is not None: - params.append(Value.I64(command.expected_version_val)) + self.bind_field(params, Value.I64(command.expected_version_val), entity, version_property.name) predicates.append(f"{self.quote_ident(version_property.column_name_val)} = {self.placeholder(len(params))}") table_name = getattr(entity, 'table_name_val', entity._name) @@ -257,7 +260,7 @@ def compile_delete(self, entity: EntityDescriptor, command: DeleteCommand) -> Co if not id_property: raise MissingIdPropertyError(entity._name) - params = [] + params = SQLBindings() 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 getattr(p, '_is_version', False)), None) @@ -265,27 +268,27 @@ def compile_delete(self, entity: EntityDescriptor, command: DeleteCommand) -> Co if not version_property: raise MissingVersionPropertyError(entity._name) if command.expected_version_val is not None: - params.append(Value.I64(-(command.expected_version_val + 1))) + self.bind_field(params, Value.I64(-(command.expected_version_val + 1)), entity, version_property.name) else: - params.append(Value.I64(-1)) + self.bind_field(params, Value.I64(-1), entity, version_property.name) - params.append(command.id) + self.bind_field(params, command.id, entity, id_property.name) predicates = [f"{self.quote_ident(id_property.column_name_val)} = {self.placeholder(len(params))}"] if command.expected_version_val is not None: - params.append(Value.I64(command.expected_version_val)) + self.bind_field(params, Value.I64(command.expected_version_val), entity, version_property.name) predicates.append(f"{self.quote_ident(version_property.column_name_val)} = {self.placeholder(len(params))}") sql = f"UPDATE {self.quote_ident(table_name)} SET {self.quote_ident(version_property.column_name_val)} = {self.placeholder(1)} WHERE {' AND '.join(predicates)}" return CompiledQuery(sql=sql, params=params) - params.append(command.id) + self.bind_field(params, command.id, entity, id_property.name) predicates = [f"{self.quote_ident(id_property.column_name_val)} = {self.placeholder(len(params))}"] if command.expected_version_val is not None: if not version_property: raise MissingVersionPropertyError(entity._name) - params.append(Value.I64(command.expected_version_val)) + self.bind_field(params, Value.I64(command.expected_version_val), entity, version_property.name) predicates.append(f"{self.quote_ident(version_property.column_name_val)} = {self.placeholder(len(params))}") sql = f"DELETE FROM {self.quote_ident(table_name)} WHERE {' AND '.join(predicates)}" @@ -303,11 +306,10 @@ def compile_recover(self, entity: EntityDescriptor, command: RecoverCommand) -> if not version_property: raise MissingVersionPropertyError(entity._name) - params = [ - Value.I64(-command.expected_version_val + 1), - command.id, - Value.I64(command.expected_version_val) - ] + params = SQLBindings() + self.bind_field(params, Value.I64(-command.expected_version_val + 1), entity, version_property.name) + self.bind_field(params, command.id, entity, id_property.name) + self.bind_field(params, Value.I64(command.expected_version_val), entity, version_property.name) table_name = getattr(entity, 'table_name_val', entity._name) sql = f"UPDATE {self.quote_ident(table_name)} SET {self.quote_ident(version_property.column_name_val)} = {self.placeholder(1)} WHERE {self.quote_ident(id_property.column_name_val)} = {self.placeholder(2)} AND {self.quote_ident(version_property.column_name_val)} = {self.placeholder(3)}" @@ -396,7 +398,39 @@ def compile_projection(self, entity: EntityDescriptor, query: SelectQuery, param else: return self.aggregate_projection(entity, query, params) + def field_log_policy(self, entity, field): + from teaql.runtime.log_privacy import credential_name + prop = entity.property_by_name(field) + if credential_name(field) or (prop and credential_name(prop.column_name_val)): + return 'credential' + if field in entity.audit_mask_fields_val: + return 'masked' + return getattr(prop, 'log_policy_val', 'unknown') + + def bind_field(self, params, value, entity, field): + if isinstance(params, SQLBindings): + params.append(value, self.field_log_policy(entity, field)) + else: + params.append(value) + + def expression_log_policy(self, entity, expr): + def columns(node): + if isinstance(node, ColumnExpr): return [node.name] + if isinstance(node, FunctionExpr): return [name for arg in node.args for name in columns(arg)] + if isinstance(node, BinaryExpr): return columns(node.left) + columns(node.right) + if isinstance(node, BetweenExpr): return columns(node.expr) + columns(node.lower) + columns(node.upper) + return [] + policies = [self.field_log_policy(entity, name) for name in columns(expr)] + # No field metadata means inherit the surrounding comparison, not plain. + return next((p for p in ('credential', 'unknown', 'masked', 'plain') if p in policies), None) + def compile_expr(self, entity: EntityDescriptor, expr: Expr, params: List[Value]) -> str: + if isinstance(params, SQLBindings): + with params.policy(self.expression_log_policy(entity, expr)): + return self._compile_expr(entity, expr, params) + return self._compile_expr(entity, expr, params) + + def _compile_expr(self, entity: EntityDescriptor, expr: Expr, params: List[Value]) -> str: if isinstance(expr, ColumnExpr): return self.column_sql(entity, expr.name) elif isinstance(expr, ValueExpr): diff --git a/src/teaql/sql/executor.py b/src/teaql/sql/executor.py index aeb07a6..4e6d8fb 100644 --- a/src/teaql/sql/executor.py +++ b/src/teaql/sql/executor.py @@ -7,7 +7,7 @@ import time import threading from array import array -from dataclasses import fields, is_dataclass +from dataclasses import fields, is_dataclass, replace from enum import Enum from teaql.data_service import ( DataServiceExecutor, QueryExecutor, MutationExecutor, @@ -32,6 +32,26 @@ _id_set_build_locks = {} _id_set_build_locks_guard = threading.RLock() + +class _QueryWithLogIntent(QueryRequest): + """Invocation-local compiler plumbing, excluded from dataclass/wire fields.""" + def __init__(self, query, trace_chain, comment, purpose, source): + super().__init__(query, trace_chain, comment, purpose) + self._log_intent_source = source + + +def _intent_bindings(compiled, request): + # Resolve each source's policy before flattening: credential detection and + # malformed-policy handling depend on the original statement, not the child. + from teaql.runtime.log_privacy import _binding_policies + inherited = getattr(request, '_log_intent_source', None) + sources = [inherited, compiled] if inherited is not None else [compiled] + return CompiledQuery('', [deepcopy(value) for source in sources for value in source.params], + parameter_log_policies=[policy for source in sources + for policy in _binding_policies(source)], + sql_origin='generated') + + def _canonical_id_set_value(value): if isinstance(value, Value): return ("Value", str(value._type_hint), _canonical_id_set_value(value.val)) @@ -160,6 +180,47 @@ def capabilities(self) -> DataServiceCapabilities: returning=False ) + def _record_statement(self, context, request, compiled, started_at, operation, + outcome, result_count=None, affected_rows=None): + """Project through context; never attach driver exceptions to diagnostics.""" + query = operation == DataServiceOperation.Query + entity = request.query.entity if query else request._data.entity + comment = request._comment if query else getattr(request, 'comment', None) + if not query and callable(comment): + comment = comment() + provider = str(self.dialect.kind()).lower() + sql_operation = 'select' if query else operation.name.lower() + metadata = ExecutionMetadata( + backend=provider, operation=operation, started_at=started_at, + ended_at=datetime.now(), execution_outcome=outcome, + parameterized_sql=compiled.sql, parameters=list(compiled.params), + parameter_log_policies=compiled.parameter_log_policies, + sql_origin=compiled.sql_origin, database_kind=self.dialect.kind(), debug_query='', + result_count=result_count, affected_rows=affected_rows, + comment=comment, purpose=request._purpose if query else None, + audit_reason=None if query else comment, + trace_chain=[ + TraceNode(kind='operation', name='query' if query else 'mutation', comment='query' if query else 'mutation'), + TraceNode(kind='request' if query else 'entity', name=entity, comment=entity), + *(request.trace_chain if query else request.trace_chain()), + TraceNode(kind='provider', name=provider, comment=provider), + TraceNode(kind='sql', name=sql_operation, comment=sql_operation), + ], + ) + if context is not None: + try: + source = getattr(request, '_log_intent_source', None) if query else None + if source is None: + context.record_metadata_log(metadata) + else: + context._record_metadata_log(metadata, intent_source=source) + except Exception: + # A broken diagnostic destination must not replace an in-flight + # driver failure, cancellation or generator close. + if outcome == 'success': + raise + return metadata + async def query_stream(self, context, request: QueryRequest, chunk_size: int): if chunk_size <= 0: raise ValueError("chunk_size must be positive") @@ -176,13 +237,40 @@ async def query_stream(self, context, request: QueryRequest, chunk_size: int): compiled = self.dialect.compile_select(entity_desc, request.query) pending = None index = 0 - async for rows in self.transport.stream_sql(compiled, chunk_size): + delivered = 0 + started_at = datetime.now() + outcome = 'cancelled' + stream = self.transport.stream_sql(compiled, chunk_size) + try: + async for rows in stream: + if pending is not None: + delivered += len(pending) + yield StreamChunk(pending, index, False) + index += 1 + pending = rows if pending is not None: - yield StreamChunk(pending, index, False) - index += 1 - pending = rows - if pending is not None: - yield StreamChunk(pending, index, True) + delivered += len(pending) + yield StreamChunk(pending, index, True) + outcome = 'success' + except (asyncio.CancelledError, GeneratorExit): + outcome = 'cancelled' + raise + except BaseException: + outcome = 'failure' + raise + finally: + try: + close = getattr(stream, 'aclose', None) + if close is not None: + await close() + except BaseException: + if outcome == 'success': + outcome = 'failure' + raise + # Preserve the error/close already being propagated. + finally: + self._record_statement(context, request, compiled, started_at, + DataServiceOperation.Query, outcome, result_count=delivered) async def query(self, context: 'UserContext', request: QueryRequest) -> QueryResult: telemetry = context.runtime_telemetry() if context is not None else None @@ -218,7 +306,10 @@ async def _query(self, context: 'UserContext', request: QueryRequest) -> QueryRe if context is not None: context.record_metadata_log(metadata) return QueryResult(rows=[], metadata=metadata) - request = QueryRequest(execution_query, request.trace_chain, request._comment, request._purpose) + source = getattr(request, '_log_intent_source', None) + request = (_QueryWithLogIntent(execution_query, request.trace_chain, request._comment, + request._purpose, source) if source is not None else + QueryRequest(execution_query, request.trace_chain, request._comment, request._purpose)) entity_desc = self.schema_provider.get_entity(request.query.entity) if not entity_desc and context: entities = context.get_resource("entities") @@ -247,11 +338,19 @@ async def _query(self, context: 'UserContext', request: QueryRequest) -> QueryRe }), lambda: self.transport.fetch_all_sql(compiled), ) - except Exception as e: - raise TransportError(e) + except BaseException as e: + self._record_statement(context, request, compiled, start, DataServiceOperation.Query, + 'cancelled' if isinstance(e, asyncio.CancelledError) else 'failure') + if isinstance(e, Exception): + raise TransportError(e) from e + raise + metadata = self._record_statement(context, request, compiled, start, + DataServiceOperation.Query, 'success', result_count=len(rows)) - await self._enhance_relations(context, rows, request) - await self._enhance_relation_aggregates(context, rows, request) + source = (_intent_bindings(compiled, request) if rows and + (request.query.relations or request.query.relation_aggregates) else None) + await self._enhance_relations(context, rows, request, source) + await self._enhance_relation_aggregates(context, rows, request, source) if retained_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 retained_order if entity_id in by_id] @@ -295,31 +394,6 @@ async def _query(self, context: 'UserContext', request: QueryRequest) -> QueryRe from teaql.core.list import SmartList facets[facet.name] = SmartList(facet_rows) - end = datetime.now() - - provider = str(self.dialect.kind()).lower() - trace_path = [ - TraceNode(kind="operation", name="query", comment="query"), - TraceNode(kind="request", name=request.query.entity, comment=request.query.entity), - *request.trace_chain, - TraceNode(kind="provider", name=provider, comment=provider), - TraceNode(kind="sql", name="select", comment="select"), - ] - metadata = ExecutionMetadata( - backend=provider, - operation=DataServiceOperation.Query, - started_at=start, - ended_at=end, - parameterized_sql=compiled.sql_with_comment(), - parameters=list(compiled.params), - result_count=len(rows), - trace_chain=trace_path, - comment=request._comment, - purpose=request._purpose, - debug_query=compiled.debug_sql(self.dialect.kind()) - ) - if context is not None: - context.record_metadata_log(metadata) return QueryResult( rows=rows, metadata=metadata, @@ -425,7 +499,8 @@ def _id_set_query_key(context, query: SelectQuery, namespace: str) -> str: _canonical_id_set_value(normalized)) return "teaql:id-set:v1:" + hashlib.sha256(repr(scope).encode("utf-8")).hexdigest() - async def _enhance_relations(self, context, parents: List[Dict[str, Any]], request: QueryRequest) -> None: + async def _enhance_relations(self, context, parents: List[Dict[str, Any]], request: QueryRequest, + intent_source=None) -> None: query = request.query if not parents or not query.relations: return @@ -469,8 +544,8 @@ async def _enhance_relations(self, context, parents: List[Dict[str, Any]], reque child_trace = [*request.trace_chain, TraceNode( kind="relation", name=f"{query.entity}.{load.name}", comment=load.name)] - children.extend((await self.query(context, QueryRequest( - probe, child_trace, request._comment, request._purpose))).rows) + children.extend((await self.query(context, _QueryWithLogIntent( + probe, child_trace, request._comment, request._purpose, intent_source))).rows) selected_plan = "bounded_probes" probe_count = len(parent_ids) else: @@ -482,8 +557,8 @@ async def _enhance_relations(self, context, parents: List[Dict[str, Any]], reque child_trace = [*request.trace_chain, TraceNode( kind="relation", name=f"{query.entity}.{load.name}", comment=load.name)] - children = (await self.query(context, QueryRequest( - child_query, child_trace, request._comment, request._purpose))).rows + children = (await self.query(context, _QueryWithLogIntent( + child_query, child_trace, request._comment, request._purpose, intent_source))).rows selected_plan = "window" if limited else "batch" probe_count = 0 for child in children: @@ -508,7 +583,8 @@ async def _enhance_relations(self, context, parents: List[Dict[str, Any]], reque relation_scope.failure(error) raise - async def _enhance_relation_aggregates(self, context, parents: List[Dict[str, Any]], request: QueryRequest) -> None: + async def _enhance_relation_aggregates(self, context, parents: List[Dict[str, Any]], request: QueryRequest, + intent_source=None) -> None: query = request.query if not parents or not query.relation_aggregates: return @@ -543,8 +619,8 @@ async def _enhance_relation_aggregates(self, context, parents: List[Dict[str, An child_trace = [*request.trace_chain, TraceNode( kind="relation", name=f"{query.entity}.{aggregate.relation_name}", comment=aggregate.relation_name)] - rows = (await self.query(context, QueryRequest( - child_query, child_trace, request._comment, request._purpose))).rows + rows = (await self.query(context, _QueryWithLogIntent( + child_query, child_trace, request._comment, request._purpose, intent_source))).rows child_desc = self.schema_provider.get_entity(relation.target_entity) foreign_property = child_desc.property_by_name(relation.foreign_key) if child_desc else None if foreign_property and foreign_property.column_name_val != relation.foreign_key: @@ -602,8 +678,13 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut result = await executor._mutate(context, request) await transaction.commit_sql() return result - except Exception: - await transaction.rollback_sql() + except BaseException: + # CancelledError is not an Exception. Release the transaction + # on cancellation too, without replacing the original failure. + try: + await transaction.rollback_sql() + except BaseException: + pass raise req_data = request._data @@ -660,6 +741,8 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut start = datetime.now() last_insert_id = None + operation = {'insert': DataServiceOperation.Insert, 'update': DataServiceOperation.Update, + 'delete': DataServiceOperation.Delete, 'recover': DataServiceOperation.Recover}[op] try: telemetry = context.runtime_telemetry() if context is not None else None affected_rows, last_insert_id = await observe_runtime_operation( @@ -670,9 +753,14 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut }), lambda: self.transport.execute_sql(compiled), ) - except Exception as e: - raise TransportError(e) - end = datetime.now() + except BaseException as e: + self._record_statement(context, request, compiled, start, operation, + 'cancelled' if isinstance(e, asyncio.CancelledError) else 'failure') + if isinstance(e, Exception): + raise TransportError(e) from e + raise + metadata = self._record_statement(context, request, compiled, start, operation, + 'success', affected_rows=affected_rows) generated_values = {} if op == "insert" and last_insert_id: @@ -710,47 +798,25 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut ) table = self.dialect.quote_ident(entity_desc.table_name_val) id_column = self.dialect.quote_ident(id_prop.column_name_val) - persisted_rows = await self.transport.fetch_all_sql( - CompiledQuery( - f"SELECT {columns} FROM {table} WHERE {id_column} = {self.dialect.placeholder(1)}", - [Value.from_any(entity_id)], - ) + readback = CompiledQuery( + f"SELECT {columns} FROM {table} WHERE {id_column} = {self.dialect.placeholder(1)}", + [Value.from_any(entity_id)], + parameter_log_policies=[self.dialect.field_log_policy(entity_desc, id_prop.name)], + sql_origin='generated', ) - if not persisted_rows: - raise TransportError(RuntimeError( - f"authoritative persisted row not found for {req_data.entity}")) + read_start = datetime.now() + persisted_rows = None + try: + persisted_rows = await self.transport.fetch_all_sql(readback) + if len(persisted_rows) != 1: + raise TransportError(RuntimeError( + f"expected one authoritative persisted row for {req_data.entity}, got {len(persisted_rows)}")) + except BaseException as error: + self._record_readback(context, readback, compiled, metadata, read_start, + persisted_rows, error) + raise persisted_record = persisted_rows[0] - provider = str(self.dialect.kind()).lower() - trace_path = [ - TraceNode(kind="operation", name="mutation", comment="mutation"), - TraceNode(kind="entity", name=req_data.entity, comment=req_data.entity), - *request.trace_chain(), - TraceNode(kind="provider", name=provider, comment=provider), - TraceNode(kind="sql", name=op, comment=op), - ] - metadata = ExecutionMetadata( - backend=provider, - operation={ - "insert": DataServiceOperation.Insert, - "update": DataServiceOperation.Update, - "delete": DataServiceOperation.Delete, - "recover": DataServiceOperation.Recover, - }[op], - started_at=start, - ended_at=end, - parameterized_sql=compiled.sql_with_comment(), - parameters=list(compiled.params), - affected_rows=affected_rows, - trace_chain=trace_path, - 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: - context.record_metadata_log(metadata) if affected_rows > 0 and context is not None: from teaql.runtime.audit import AuditFieldChange, MutationAuditKind, RawAuditEvent if isinstance(req_data, InsertCommand): @@ -786,6 +852,25 @@ async def _mutate(self, context: 'UserContext', request: MutationRequest) -> Mut persisted_record=persisted_record, ) + def _record_readback(self, context, readback, source, write_metadata, started_at, rows, error): + if context is None: + return + # A driver returning zero/multiple rows succeeded as SQL; validation of + # the authoritative snapshot is a separate business failure. + outcome = ('success' if rows is not None else 'cancelled' + if isinstance(error, asyncio.CancelledError) else 'failure') + metadata = replace(write_metadata, operation=DataServiceOperation.Query, + started_at=started_at, ended_at=datetime.now(), execution_outcome=outcome, + parameterized_sql=readback.sql, parameters=list(readback.params), + parameter_log_policies=readback.parameter_log_policies, sql_origin=readback.sql_origin, + affected_rows=None, result_count=len(rows) if rows is not None else None, + trace_chain=[*write_metadata.trace_chain, TraceNode(kind='sql', name='readback', comment='readback')]) + try: + context._record_metadata_log(metadata, intent_source=source) + except BaseException: + # An in-flight readback error must survive a diagnostic sink failure. + pass + async def next_id(self, entity: str) -> int: await self.transport.execute_sql(CompiledQuery( "CREATE TABLE IF NOT EXISTS teaql_id_space (" diff --git a/src/teaql/sql/types.py b/src/teaql/sql/types.py index 2e2d1a1..abf93e2 100644 --- a/src/teaql/sql/types.py +++ b/src/teaql/sql/types.py @@ -4,6 +4,7 @@ from datetime import date, datetime, timezone from decimal import Decimal import json +from contextlib import contextmanager from teaql.core.value import Value, DataType, Timestamp class DatabaseKind(Enum): @@ -11,11 +12,39 @@ class DatabaseKind(Enum): Sqlite = auto() MySql = auto() +class SQLBindings(list): + """Compiler-owned provenance; never supplied by incoming query payloads.""" + def __init__(self): + super().__init__() + self.policies = [] + self.current_policy = 'unknown' + self.trusted = True + + def append(self, value, policy=None): + super().append(value) + self.policies.append(policy or self.current_policy) + + @contextmanager + def policy(self, policy): + previous = self.current_policy + self.current_policy = policy or previous + try: + yield + finally: + self.current_policy = previous + @dataclass class CompiledQuery: sql: str params: List[Value] comment: Optional[str] = None + parameter_log_policies: Optional[List[str]] = None + sql_origin: Optional[str] = None + + def __post_init__(self): + if isinstance(self.params, SQLBindings): + self.parameter_log_policies = list(self.params.policies) + self.sql_origin = 'generated' if self.params.trusted else None def sql_with_comment(self) -> str: if self.comment: @@ -28,7 +57,24 @@ def debug_sql(self, dialect: DatabaseKind) -> str: return _replace_numbered_placeholders(self.sql_with_comment(), self.params, dialect) return _replace_positional_placeholders(self.sql_with_comment(), self.params, dialect) -def _replace_numbered_placeholders(sql: str, params: List[Value], dialect: DatabaseKind) -> str: +def render_log_sql(sql, params, dialect, literal): + """Use the normal dialect scanner, with complete bind accounting for logs.""" + if not sql.strip(): + raise ValueError('Missing SQL template') + used = set() + def checked(index): + if index < 0 or index >= len(params): + raise ValueError('SQL bind count mismatch') + used.add(index) + return literal(index) + scanner = (_replace_numbered_placeholders if dialect == DatabaseKind.PostgreSql + else _replace_positional_placeholders) + rendered = scanner(sql, params, dialect, checked) + if len(used) != len(params): + raise ValueError('Unused SQL bindings') + return rendered + +def _replace_numbered_placeholders(sql: str, params: List[Value], dialect: DatabaseKind, literal=None) -> str: output, index, state = [], 0, "sql" while index < len(sql): char = sql[index] @@ -56,16 +102,19 @@ def _replace_numbered_placeholders(sql: str, params: List[Value], dialect: Datab while end < len(sql) and sql[end].isdigit(): end += 1 parameter_index = int(sql[index + 1:end]) - 1 - output.append(_sql_literal(params[parameter_index], dialect) + output.append(literal(parameter_index) if literal else + _sql_literal(params[parameter_index], dialect) if 0 <= parameter_index < len(params) else sql[index:end]) index = end continue else: output.append(char) index += 1 + if literal and state not in ('sql', 'line'): + raise ValueError('Unclosed SQL token') return "".join(output) -def _replace_positional_placeholders(sql: str, params: List[Value], dialect: DatabaseKind) -> str: +def _replace_positional_placeholders(sql: str, params: List[Value], dialect: DatabaseKind, literal=None) -> str: output, index, parameter_index, state = [], 0, 0, "sql" while index < len(sql): char = sql[index] @@ -76,6 +125,8 @@ def _replace_positional_placeholders(sql: str, params: List[Value], dialect: Dat elif state == "sql" and char == '"': output.append(char) state = "double_quote" + elif state == "sql" and char == '`': + output.append(char); state = "backtick" elif state == "sql" and char == "-" and next_char == "-": output.extend((char, next_char)); index += 1; state = "line_comment" elif state == "sql" and char == "/" and next_char == "*": @@ -96,17 +147,23 @@ def _replace_positional_placeholders(sql: str, params: List[Value], dialect: Dat elif state == "line_comment": output.append(char) if char in "\r\n": state = "sql" + elif state == "backtick": + output.append(char) + if char == '`' and next_char == '`': output.append('`'); index += 1 + elif char == '`': state = 'sql' elif state == "block_comment": output.append(char) if char == "*" and next_char == "/": output.append("/"); index += 1; state = "sql" - elif (char == "?" or (dialect == DatabaseKind.MySql and char == "%" and next_char == "s")) and parameter_index < len(params): - output.append(_sql_literal(params[parameter_index], dialect)) + elif (char == "?" or (dialect == DatabaseKind.MySql and char == "%" and next_char == "s")) and (literal or parameter_index < len(params)): + output.append(literal(parameter_index) if literal else _sql_literal(params[parameter_index], dialect)) parameter_index += 1 if char == "%": index += 1 else: output.append(char) index += 1 + if literal and state not in ('sql', 'line_comment'): + raise ValueError('Unclosed SQL token') return "".join(output) def _sql_literal(value: Value, dialect: DatabaseKind) -> str: diff --git a/test-vectors/masking-v1.tsv b/test-vectors/masking-v1.tsv new file mode 100644 index 0000000..476f5a8 --- /dev/null +++ b/test-vectors/masking-v1.tsv @@ -0,0 +1,11 @@ +id input expected +empty +short Ada *** +digits 12345678 ******** +boundary ABCDEFGH AB****GH +long Riverside Ri*****de +quote O'Reilly O'****ly +emoji 😀😀1234😀😀 😀😀****😀😀 +non_ascii_digits 12345678 12****78 +combining éabcdef é****ef +cjk 甲乙丙丁戊己庚辛 甲乙****庚辛 diff --git a/tests/core/test_context.py b/tests/core/test_context.py index de3be57..11e3c4c 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -40,7 +40,9 @@ class MockMetadata: assert len(context.sql_logs()) == 1 assert context.sql_logs()[0].debug_sql == "[REDACTED SQL; NOT REPLAYABLE]" assert output[0].endswith("\nDebug SQL: [REDACTED SQL; NOT REPLAYABLE]") - assert "Parameterized SQL:" in output[0] + # #25/#58: text diagnostics show expanded SQL, not placeholders plus a list. + assert "Parameterized SQL:" not in output[0] + assert "SQL omission reason: unsupported-or-mismatched-bindings" in output[0] # SQL logs class MockQuery: diff --git a/tests/runtime/test_log_privacy.py b/tests/runtime/test_log_privacy.py index 83822d2..3d91cce 100644 --- a/tests/runtime/test_log_privacy.py +++ b/tests/runtime/test_log_privacy.py @@ -54,7 +54,7 @@ def test_exact_opt_in_warns_and_preserves_noncredential_debug(monkeypatch, caplo from teaql.runtime.log_privacy import _warn_plaintext _warn_plaintext.cache_clear() monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) - raw = entry() + raw = replace(entry(), parameter_log_policies=['masked']) output = [] TextDiagnosticSqlLogSink(output.append).write(raw) assert raw.params[0] in output[0] @@ -86,7 +86,9 @@ def test_audit_field_policy_and_credential_override(monkeypatch, enabled): AuditFieldChange("password", "old-password-secret", "new-password-secret"), ), ("change new-password-secret",)) safe = event.safe(["email"], None) - assert safe.fields[0].value == ("new-private-email" if enabled else REDACTED) + # #58 restores the existing business mask; credentials below remain fully hidden. + assert safe.fields[0].value == ("new-private-email" if enabled else "ne*************il") + assert safe.fields[0].masked == (not enabled) assert safe.fields[1].value == REDACTED assert "new-password-secret" not in repr(safe) assert event.changes[1].new_value == "new-password-secret" diff --git a/tests/runtime/test_masking_contract.py b/tests/runtime/test_masking_contract.py new file mode 100644 index 0000000..03f373f --- /dev/null +++ b/tests/runtime/test_masking_contract.py @@ -0,0 +1,66 @@ +"""Behavior contracts: never replace expected output just to match a new implementation.""" +from datetime import datetime, timedelta +from dataclasses import replace +from pathlib import Path +import pytest +from teaql.runtime.audit import _mask +from teaql.runtime.context import SqlLogEntry, SqlLogOperation, TextDiagnosticSqlLogSink +from teaql.runtime.log_privacy import PLAINTEXT_ENV, PLAINTEXT_ACK + +CASES = [("", ""), ("Ada", "***"), ("12345678", "********"), + ("ABCDEFGH", "AB****GH"), ("Riverside", "Ri*****de"), ("O'Reilly", "O'****ly")] + +GOLDEN = [line.split("\t") for line in + (Path(__file__).parents[2] / "test-vectors/masking-v1.tsv").read_text().splitlines()[1:]] + +@pytest.mark.parametrize("case,raw,expected", GOLDEN) +def test_mask_golden(case, raw, expected, monkeypatch): + from teaql.runtime.audit import AuditFieldChange, MutationAuditKind, RawAuditEvent + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + assert _mask(raw) == expected, case + event = RawAuditEvent(MutationAuditKind.UPDATED, "Customer", 1, + (AuditFieldChange("name", None, raw),)) + safe = event.safe(["name"], None) + assert safe.fields[0].masked + assert safe.fields[0].value == expected, case + assert event.changes[0].new_value == raw + +def entry(value): + sql = "UPDATE customer SET name = '" + value.replace("'", "''") + "'" + return SqlLogEntry(SqlLogOperation.Update, "what: edit customer", "why: verify mask contract", + None, [], "UPDATE customer SET name = ?", [value], sql, sql, + datetime.now(), datetime.now(), timedelta(microseconds=1), None, None, 1, "1 row affected") + +@pytest.mark.parametrize("raw,masked", CASES) +def test_mask_contract_legacy_algorithm(raw, masked): + assert _mask(raw) == masked + +@pytest.mark.parametrize("raw,masked", CASES) +def test_mask_contract_expanded_sql(monkeypatch, raw, masked): + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + source = entry(raw) + output = [] + TextDiagnosticSqlLogSink(output.append).write(source) + log = "\n".join(output) + assert source.params == [raw] + # Unknown parameter provenance must not expose a prefix/suffix. + assert "name = '" in log + if len(raw) >= 8 and not masked.startswith("*"): + assert masked.replace("'", "''") not in log + assert "masked" in log.lower() + assert "name = ?" not in log + assert "[REDACTED SQL" not in log + if raw: + assert "'" + raw.replace("'", "''") + "'" not in log + +def test_mask_contract_debug_provenance(monkeypatch): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + output = [] + sink = TextDiagnosticSqlLogSink(output.append) + sink.write(replace(entry("Riverside"), parameter_log_policies=['masked'])) + sink.write(replace(entry("Riverside"), parameter_log_policies=['masked'])) + assert len(output) == 2 + for log in output: + assert "'Riverside'" in log + assert "DEBUG" in log.upper() + assert "PLAINTEXT" in log.upper() diff --git a/tests/runtime/test_relation_masking.py b/tests/runtime/test_relation_masking.py new file mode 100644 index 0000000..a99fa86 --- /dev/null +++ b/tests/runtime/test_relation_masking.py @@ -0,0 +1,150 @@ +"""Actual SQLite relation plans must inherit intent redactions, not SQL binds.""" +from types import SimpleNamespace + +import aiosqlite +import pytest + +from teaql.core.expr import Expr +from teaql.core.meta import EntityDescriptor, PropertyDescriptor, RelationDescriptor +from teaql.core.query import SelectQuery, RelationAggregate +from teaql.core.value import DataType +from teaql.data_service import QueryRequest +from teaql.provider.sqlite import SimpleSchemaProvider +from teaql.provider.sqlite.dialect import SqliteDialect +from teaql.provider.sqlite.transport import SqliteTransport +from teaql.runtime import RuntimeModule +from teaql.runtime.context import TextDiagnosticSqlLogSink +from teaql.runtime.log_privacy import PLAINTEXT_ENV, PLAINTEXT_ACK, sql_log_projection +from teaql.sql.executor import SqlDataServiceExecutor, TransportError + + +class CaptureTransport(SqliteTransport): + def __init__(self, path): + super().__init__(path) + self.reads = [] + self.failing_table = None + self.failure = RuntimeError('DRIVER-CANARY') + + async def fetch_all_sql(self, compiled): + self.reads.append(compiled) + if self.failing_table and self.failing_table in compiled.sql: + raise self.failure + return await super().fetch_all_sql(compiled) + + +class CaptureExecutor(SqlDataServiceExecutor): + def __init__(self, transport, provider): + super().__init__(SqliteDialect(), transport, provider) + self.dispatched = [] + + async def query(self, context, request): + self.dispatched.append(request.query.entity) + return await super().query(context, request) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('shape', ['batch', 'probe', 'window', 'aggregate', 'nested']) +@pytest.mark.parametrize('debug', [False, True]) +@pytest.mark.parametrize('failure', [False, True]) +async def test_derived_relation_intent(tmp_path, monkeypatch, shape, debug, failure): + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + provider = SimpleSchemaProvider() + module = RuntimeModule.new() + for name, fields in [('Customer', ['name', 'password']), + ('Order', ['customer_id', 'name']), ('Line', ['order_id'])]: + entity = EntityDescriptor(name).table_name(name.lower() + '_data') + entity.property(PropertyDescriptor('id', DataType.I64).is_id().log_policy('plain')) + entity.property(PropertyDescriptor('version', DataType.I64).is_version().log_policy('plain')) + for field in fields: + entity.property(PropertyDescriptor(field, DataType.I64 if field.endswith('_id') + else DataType.Text).log_policy('plain')) + entity.audit_mask_fields(['name']) + if name == 'Customer': + entity.relation(RelationDescriptor('orders', 'Order').foreign('customer_id').many()) + elif name == 'Order': + entity.relation(RelationDescriptor('lines', 'Line').foreign('order_id').many()) + provider.register_entity(entity) + module.entity(entity) + path = str(tmp_path / 'relations.db') + async with aiosqlite.connect(path) as db: + await db.executescript(''' + CREATE TABLE customer_data(id INTEGER PRIMARY KEY, version INTEGER, name TEXT, password TEXT); + CREATE TABLE order_data(id INTEGER PRIMARY KEY, version INTEGER, customer_id INTEGER, name TEXT); + CREATE TABLE line_data(id INTEGER PRIMARY KEY, version INTEGER, order_id INTEGER); + ''') + await db.execute('INSERT INTO customer_data VALUES (1,1,?,?)', ('Riverside', 'PASSWORD-CANARY')) + await db.execute('INSERT INTO order_data VALUES (1,1,1,?)', ('Lakeside',)) + await db.execute('INSERT INTO line_data VALUES (1,1,1)') + await db.commit() + transport = CaptureTransport(path) + service = CaptureExecutor(transport, provider) + context = module.into_context() + logs, output = [], [] + sink = TextDiagnosticSqlLogSink(output.append) + def capture(entry): + logs.append(entry) + sink.write(entry) + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=capture)) + if debug: + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + if failure: + transport.failing_table = 'line_data' if shape == 'nested' else 'order_data' + query = (SelectQuery('Customer').project('id', 'name') + .filter(Expr.eq('name', 'Riverside')).and_filter(Expr.eq('password', 'PASSWORD-CANARY')).limit(1)) + child = SelectQuery('Order').project('id', 'name') + if shape == 'probe': + child.limit(1).top_n_probe_parent_threshold(32) + if shape == 'window': + child.limit(1).top_n_probe_parent_threshold(0) + if shape == 'nested': + child.filter(Expr.eq('name', 'Lakeside')).relation_query('lines', SelectQuery('Line').project('id').limit(2)) + if shape == 'aggregate': + query.relation_aggregates.append(RelationAggregate('orders', 'record_count', SelectQuery('Order').count('n'), True)) + elif shape == 'batch': + query.relation('orders') + else: + query.relation_query('orders', child) + request = (QueryRequest(query).comment('what: load Riverside PASSWORD-CANARY Lakeside graph') + .purpose('why: verify inherited intent')) + if failure: + with pytest.raises(TransportError) as caught: + await service.query(context, request) + assert caught.value.error is transport.failure + else: + result = await service.query(context, request) + assert len(result.rows) == 1 + if shape == 'aggregate': + assert result.rows[0]['record_count'] == 1 + elif shape == 'batch': + # A relation without a nested projection loads only its linking key. + assert result.rows[0]['orders'] == [{'customer_id': 1}] + else: + assert result.rows[0]['orders'][0]['id'] == 1 + if shape == 'nested': + assert result.rows[0]['orders'][0]['lines'][0]['id'] == 1 + assert service.dispatched == (['Customer', 'Order', 'Line'] if shape == 'nested' else ['Customer', 'Order']) + assert len(logs) == len(service.dispatched) + entry = logs[-1] + assert entry.execution_outcome == ('failure' if failure else 'success') + assert entry.comment.startswith('what: load') + assert entry.purpose == 'why: verify inherited intent' + assert ('Riverside' in entry.comment) == debug + assert 'PASSWORD-CANARY' not in repr(logs) + '\n'.join(output) + repr(context.sql_logs()) + if not debug: + assert 'Riverside' not in repr(logs) + '\n'.join(output) + if shape == 'nested': + assert 'Lakeside' not in entry.comment + assert len(entry.params) == len(transport.reads[-1].params) + assert 'PASSWORD-CANARY' in [v.val for v in transport.reads[0].params] + if shape == 'batch': + assert ' IN (' in entry.sql + elif shape == 'window': + assert 'ROW_NUMBER() OVER' in entry.sql + elif shape == 'probe': + assert ' IN (' not in entry.sql and 'ROW_NUMBER' not in entry.sql + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + assert 'Riverside' not in repr(sql_log_projection(entry)) + transport.failing_table = None + await service.query(context, QueryRequest(SelectQuery('Customer').limit(1)) + .comment('what: independent Riverside').purpose('why: source isolation')) + assert logs[-1].comment == 'what: independent Riverside' diff --git a/tests/runtime/test_sql_mask_lifecycle.py b/tests/runtime/test_sql_mask_lifecycle.py new file mode 100644 index 0000000..2b83b58 --- /dev/null +++ b/tests/runtime/test_sql_mask_lifecycle.py @@ -0,0 +1,421 @@ +import asyncio +from types import SimpleNamespace + +import pytest + +from teaql.core.expr import Expr +from teaql.core.meta import EntityDescriptor, PropertyDescriptor +from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest, TraceNode +from teaql.core.query import SelectQuery +from teaql.core.value import DataType, Value +from teaql.data_service import QueryRequest +from teaql.provider.sqlite import SimpleSchemaProvider +from teaql.provider.sqlite.dialect import SqliteDialect +from teaql.runtime import RuntimeModule +from teaql.runtime.context import TextDiagnosticSqlLogSink +from teaql.runtime.log_privacy import PLAINTEXT_ENV, PLAINTEXT_ACK +from teaql.sql.executor import (SqlDataServiceExecutor, SqlTransport, TransportError, + SqlTransactionTransport, SqlTransactionTransportTx) + + +class FaultTransport(SqlTransport): + def __init__(self, batches=0, failure=None): + self.batches, self.failure = batches, failure + self.closed = False + self.params = [] + + async def fetch_all_sql(self, query): + self.params = query.params + if self.failure: + raise self.failure + return [] + + async def execute_sql(self, query): + self.params = query.params + if self.failure: + raise self.failure + return 0, None + + async def stream_sql(self, query, chunk_size): + self.params = query.params + try: + for _ in range(self.batches): + yield [{'display_name': 'Riverside'}] + if self.failure: + raise self.failure + finally: + self.closed = True + + +class DomainStatementExecutor(SqlDataServiceExecutor): + # Isolate domain-statement failures from unrelated ID-space initialization. + async def next_id(self, entity): + return 1 + + async def ensure_id_floor(self, entity, floor): + pass + + +@pytest.fixture +def fixture(monkeypatch): + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + entity = EntityDescriptor('Customer').table_name('customer_data') + entity.property(PropertyDescriptor('id', DataType.I64).is_id()) + entity.property(PropertyDescriptor('version', DataType.I64).is_version()) + for field in ['display_name', 'public_address', 'password_hash']: + entity.property(PropertyDescriptor(field, DataType.Text).log_policy('plain')) + entity.audit_mask_fields(['display_name', 'password_hash']) + provider = SimpleSchemaProvider() + provider.register_entity(entity) + context = RuntimeModule.new().entity(entity).into_context() + output, entries = [], [] + sink = TextDiagnosticSqlLogSink(output.append) + def capture(entry): + entries.append(entry) + sink.write(entry) + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=capture)) + return context, provider, output, entries + + +def request(): + query = (SelectQuery('Customer').filter(Expr.eq('display_name', 'Riverside')) + .and_filter(Expr.eq('public_address', '1 Runtime Road')) + .and_filter(Expr.eq('password_hash', 'PASSWORD-CANARY')).limit(5)) + return QueryRequest(query).comment('what: inspect customers').purpose('why: lifecycle test') + + +def assert_log(fixture, outcome, count=None, debug=False): + context, _, output, entries = fixture + assert len(entries) == 1 + assert entries[0].execution_outcome == outcome + assert entries[0].result_count == count + text = '\n'.join(output) + assert f'outcome={outcome}' in text + assert 'PASSWORD-CANARY' not in text and 'DRIVER-CANARY' not in text + assert 'PASSWORD-CANARY' not in repr(context.sql_logs()) + if not debug: + assert 'Riverside' not in text + assert entries[0].debug_sql + + +@pytest.mark.asyncio +@pytest.mark.parametrize('operation', ['query', 'insert', 'update', 'delete']) +async def test_failed_statement_logs_and_preserves_error(fixture, operation): + context, provider, output, entries = fixture + failure = RuntimeError('DRIVER-CANARY Riverside PASSWORD-CANARY') + transport = FaultTransport(failure=failure) + executor = DomainStatementExecutor(SqliteDialect(), transport, provider) + with pytest.raises(TransportError) as caught: + if operation == 'query': + await executor.query(context, request()) + else: + commands = { + 'insert': InsertCommand('Customer').value('display_name', 'Riverside'), + 'update': UpdateCommand('Customer', Value.I64(1)).expected_version(1).value('display_name', 'Riverside'), + 'delete': DeleteCommand('Customer', Value.I64(1)).expected_version(1), + } + cmd = commands[operation] + cmd.trace_chain = [TraceNode(comment='what: lifecycle mutation')] + await executor.mutate(context, MutationRequest(cmd)) + assert caught.value.error is failure + assert_log(fixture, 'failure') + assert entries[0].affected_rows is None + if operation == 'query': + for text in ['Ri*****de', '1 Runtime Road', 'what: inspect customers', 'why: lifecycle test']: + assert text in output[0] + assert 'Riverside' in [p.val for p in transport.params] + + +@pytest.mark.asyncio +@pytest.mark.parametrize('batches,fail,count', [(3,False,3),(0,False,0),(3,True,2),(1,True,0)]) +async def test_stream_completion_and_failure(fixture, batches, fail, count): + context, provider, output, _ = fixture + failure = RuntimeError('DRIVER-CANARY Riverside PASSWORD-CANARY') if fail else None + transport = FaultTransport(batches, failure) + executor = SqlDataServiceExecutor(SqliteDialect(), transport, provider) + received = [] + try: + async for chunk in executor.query_stream(context, request(), 1): + received.extend(chunk.rows) + except RuntimeError as error: + assert error is failure + else: + assert not fail + assert len(received) == count and transport.closed + assert all(row['display_name'] == 'Riverside' for row in received) + assert_log(fixture, 'failure' if fail else 'success', count) + assert '1 Runtime Road' in output[0] and 'Ri*****de' in output[0] + + +@pytest.mark.asyncio +@pytest.mark.parametrize('batches', [1,3]) +async def test_explicit_stream_close_releases_inner_generator(fixture, batches): + context, provider, _, _ = fixture + transport = FaultTransport(batches) + stream = SqlDataServiceExecutor(SqliteDialect(), transport, provider).query_stream(context, request(), 1) + await stream.__anext__() + await stream.aclose() + assert transport.closed # must not wait for GC/event-loop finalization + assert_log(fixture, 'cancelled', 1) + + +@pytest.mark.asyncio +async def test_cancelled_stream(fixture): + context, provider, _, _ = fixture + failure = asyncio.CancelledError() + transport = FaultTransport(failure=failure) + executor = SqlDataServiceExecutor(SqliteDialect(), transport, provider) + with pytest.raises(asyncio.CancelledError) as caught: + async for _ in executor.query_stream(context, request(), 1): + pass + assert caught.value is failure and transport.closed + assert_log(fixture, 'cancelled', 0) + + +@pytest.mark.asyncio +async def test_debug_stream_still_masks_credentials(fixture, monkeypatch): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + context, provider, output, _ = fixture + executor = SqlDataServiceExecutor(SqliteDialect(), FaultTransport(), provider) + async for _ in executor.query_stream(context, request(), 1): + pass + assert_log(fixture, 'success', 0, debug=True) + assert 'Riverside' in output[0] and 'DEBUG' in output[0] and 'EXPLICIT OPT-IN' in output[0] + + +@pytest.mark.asyncio +async def test_disabled_stream_still_executes(fixture): + context, provider, output, entries = fixture + context.disable_select_sql_log() + transport = FaultTransport(2) + executor = SqlDataServiceExecutor(SqliteDialect(), transport, provider) + rows = [chunk async for chunk in executor.query_stream(context, request(), 1)] + assert len(rows) == 2 and transport.closed + assert not entries and not output and not context.sql_logs() + + +@pytest.mark.asyncio +@pytest.mark.parametrize('streaming', [False, True]) +async def test_broken_sink_cannot_replace_driver_failure(fixture, streaming): + context, provider, _, _ = fixture + def broken(entry): + raise ValueError('diagnostic destination unavailable') + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=broken)) + failure = RuntimeError('DRIVER-CANARY') + transport = FaultTransport(failure=failure) + executor = SqlDataServiceExecutor(SqliteDialect(), transport, provider) + if streaming: + with pytest.raises(RuntimeError) as caught: + async for _ in executor.query_stream(context, request(), 1): + pass + assert caught.value is failure and transport.closed + else: + with pytest.raises(TransportError) as caught: + await executor.query(context, request()) + assert caught.value.error is failure + assert len(context.sql_logs()) == 1 + + +@pytest.mark.asyncio +async def test_query_cancellation_preserved(fixture): + context, provider, _, _ = fixture + failure = asyncio.CancelledError() + executor = SqlDataServiceExecutor(SqliteDialect(), FaultTransport(failure=failure), provider) + with pytest.raises(asyncio.CancelledError) as caught: + await executor.query(context, request()) + assert caught.value is failure + assert_log(fixture, 'cancelled') + + +@pytest.mark.asyncio +async def test_consumer_athrow_closes_stream(fixture): + context, provider, _, _ = fixture + transport = FaultTransport(3) + stream = SqlDataServiceExecutor(SqliteDialect(), transport, provider).query_stream(context, request(), 1) + await stream.__anext__() + failure = RuntimeError('consumer failure') + with pytest.raises(RuntimeError) as caught: + await stream.athrow(failure) + assert caught.value is failure and transport.closed + assert_log(fixture, 'failure', 1) + + +@pytest.mark.asyncio +async def test_cancelled_mutation_rolls_back_transaction(fixture): + context, provider, _, _ = fixture + failure = asyncio.CancelledError() + class Transaction(FaultTransport, SqlTransactionTransportTx): + committed = False + rolled_back = False + async def commit_sql(self): + self.committed = True + async def rollback_sql(self): + self.rolled_back = True + tx = Transaction(failure=failure) + class Transport(FaultTransport, SqlTransactionTransport): + async def begin_sql(self): + return tx + executor = SqlDataServiceExecutor(SqliteDialect(), Transport(), provider) + cmd = DeleteCommand('Customer', Value.I64(1)).expected_version(1) + cmd.trace_chain = [TraceNode(comment='what: cancel mutation')] + with pytest.raises(asyncio.CancelledError) as caught: + await executor.mutate(context, MutationRequest(cmd)) + assert caught.value is failure + assert tx.rolled_back and not tx.committed + assert_log(fixture, 'cancelled') + + +@pytest.mark.asyncio +async def test_generated_mutation_string_comment_is_preserved(fixture): + context, provider, _, entries = fixture + executor = DomainStatementExecutor(SqliteDialect(), FaultTransport(), provider) + command = InsertCommand('Customer').value('display_name', 'Riverside') + mutation = MutationRequest(command) + mutation.comment = 'what: generated audited save' + await executor.mutate(context, mutation) + assert entries[0].comment == mutation.comment + assert entries[0].audit_reason == mutation.comment + assert_log(fixture, 'success') + + +class ReadbackTransaction(FaultTransport, SqlTransactionTransportTx): + def __init__(self, rows=None, failure=None): + super().__init__(failure=failure) + self.rows = rows or [] + self.writes = self.reads = self.commits = self.rollbacks = 0 + + async def execute_sql(self, query): + self.params = query.params + self.writes += 1 + return 1, None + + async def fetch_all_sql(self, query): + self.reads += 1 + if self.failure: + raise self.failure + return self.rows + + async def commit_sql(self): + self.commits += 1 + + async def rollback_sql(self): + self.rollbacks += 1 + + +def readback_request(): + command = (UpdateCommand('Customer', Value.I64(1)).expected_version(1) + .value('display_name', 'Riverside').value('password_hash', 'PASSWORD-CANARY')) + command.trace_chain = [TraceNode(comment='what: update Riverside PASSWORD-CANARY')] + return MutationRequest(command) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('explicit', [False, True]) +@pytest.mark.parametrize('mode', ['error', 'cancelled', 'empty', 'multiple', 'success']) +async def test_readback_independent_diagnostic(fixture, explicit, mode): + from dataclasses import asdict + context, provider, output, entries = fixture + failure = (asyncio.CancelledError() if mode == 'cancelled' + else RuntimeError('DRIVER-CANARY') if mode == 'error' else None) + row = {'id': 1, 'version': 2, 'display_name': 'Riverside'} + rows = [row] * (2 if mode == 'multiple' else 1 if mode == 'success' else 0) + tx = ReadbackTransaction(rows, failure) + class Transport(FaultTransport, SqlTransactionTransport): + async def begin_sql(self): + return tx + executor = SqlDataServiceExecutor(SqliteDialect(), Transport(), provider) + target = await executor.begin(context) if explicit else executor + if mode == 'success': + result = await target.mutate(context, readback_request()) + assert result.persisted_record['display_name'] == 'Riverside' + if explicit: + assert tx.commits == 0 + await target.commit(context) + assert tx.commits == 1 and tx.rollbacks == 0 + assert len(entries) == 1 + else: + with pytest.raises(BaseException) as caught: + await target.mutate(context, readback_request()) + if failure: + assert caught.value is failure + else: + assert isinstance(caught.value, TransportError) + if explicit: + assert tx.rollbacks == 0 + await target.rollback(context) + assert tx.rollbacks == 1 and tx.commits == 0 + assert len(entries) == 2 + read = entries[1] + assert read.execution_outcome == ('cancelled' if mode == 'cancelled' else 'failure' if mode == 'error' else 'success') + assert read.result_count == (None if failure else len(rows)) + assert read.affected_rows is None + assert read.audit_reason and 'what: update' in read.audit_reason + assert 'SELECT' in read.debug_sql + assert entries[0].execution_outcome == 'success' and entries[0].affected_rows == 1 + assert tx.writes == tx.reads == 1 + assert 'Riverside' in [value.val for value in tx.params] + for secret in ['Riverside', 'PASSWORD-CANARY', 'DRIVER-CANARY']: + assert secret not in repr([asdict(entry) for entry in entries]) + assert secret not in '\n'.join(output) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('mode', ['debug', 'disabled', 'broken-sink']) +async def test_readback_sink_and_debug_boundaries(fixture, monkeypatch, mode): + from dataclasses import replace + from teaql.runtime.log_privacy import sql_log_projection + context, provider, output, entries = fixture + if mode == 'debug': + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + if mode == 'disabled': + context.disable_sql_log() + if mode == 'broken-sink': + def broken(entry): + if 'SELECT' in entry.sql: + raise ValueError('SINK-FAILURE') + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=broken)) + failure = RuntimeError('DRIVER-CANARY') + tx = ReadbackTransaction(failure=failure) + executor = SqlDataServiceExecutor(SqliteDialect(), tx, provider) + with pytest.raises(RuntimeError) as caught: + await executor.mutate(context, readback_request()) + assert caught.value is failure + if mode == 'disabled': + assert not entries and not context.sql_logs() + else: + assert len(context.sql_logs()) == 2 + if mode == 'debug': + assert all('Riverside' in entry.audit_reason for entry in entries) + assert all('EXPLICIT OPT-IN' in entry.debug_sql for entry in entries) + assert 'PASSWORD-CANARY' not in repr(context.sql_logs()) + monkeypatch.delenv(PLAINTEXT_ENV) + for entry in entries: + safe = sql_log_projection(entry) + assert 'what: update' in safe.audit_reason + assert 'Riverside' not in repr(safe) + assert 'Riverside' not in repr(sql_log_projection(replace(entry))) + entry.audit_reason += ' modified by sink' + assert 'Riverside' not in repr(sql_log_projection(entry)) + + +@pytest.mark.asyncio +async def test_partial_transaction_keeps_prior_write_and_stops_after_readback(fixture): + context, provider, _, entries = fixture + tx = ReadbackTransaction(rows=[{'id':1, 'version':2, 'display_name':'Riverside'}]) + class Transport(FaultTransport, SqlTransactionTransport): + async def begin_sql(self): + return tx + context.insert_resource('dataService', SqlDataServiceExecutor(SqliteDialect(), Transport(), provider)) + failure = RuntimeError('SECOND-READBACK-FAILURE') + async def work(): + service = context.require_resource('dataService') + await service.mutate(context, readback_request()) + tx.failure = failure + await service.mutate(context, readback_request()) + await service.mutate(context, readback_request()) + with pytest.raises(RuntimeError) as caught: + await context.execute_graph_save(work) + assert caught.value is failure + assert [entry.execution_outcome for entry in entries] == ['success','success','failure'] + assert tx.writes == tx.reads == 2 and tx.rollbacks == 1 and tx.commits == 0 + assert 'Riverside' not in repr(context.sql_logs()) diff --git a/tests/runtime/test_sql_masking_policy.py b/tests/runtime/test_sql_masking_policy.py new file mode 100644 index 0000000..9c3ca75 --- /dev/null +++ b/tests/runtime/test_sql_masking_policy.py @@ -0,0 +1,223 @@ +from dataclasses import replace +from datetime import datetime, timedelta +from types import SimpleNamespace + +import pytest + +from teaql.core.expr import Expr +from teaql.core.meta import EntityDescriptor, PropertyDescriptor +from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest, TraceNode +from teaql.core.query import SelectQuery +from teaql.core.value import DataType, Value +from teaql.data_service import QueryRequest +from teaql.provider.sqlite import create_sqlite_service, SimpleSchemaProvider +from teaql.provider.postgres.dialect import PostgresDialect +from teaql.runtime import RuntimeModule +from teaql.runtime.context import SqlLogEntry, SqlLogOperation, TextDiagnosticSqlLogSink +from teaql.runtime.log_privacy import PLAINTEXT_ENV, PLAINTEXT_ACK, sql_log_projection +from teaql.sql.types import DatabaseKind + + +@pytest.fixture(autouse=True) +def safe_environment(monkeypatch): + monkeypatch.delenv(PLAINTEXT_ENV, raising=False) + + +def entry(sql, params, **kwargs): + return SqlLogEntry(SqlLogOperation.Select, 'what: read customers', 'why: test policies', None, + [], sql, params, '', '', datetime.now(), datetime.now(), timedelta(microseconds=1), + 1, None, None, '1 rows returned', **kwargs) + + +def test_mixed_policies_repeated_binds_and_projection_copy(): + raw = entry('SELECT $1, $2, $1', [Value.Text("O'Reilly"), Value.Bool(True)], + database_kind=DatabaseKind.PostgreSql, parameter_log_policies=['masked', 'plain']) + safe = sql_log_projection(raw) + assert "'O''****ly' /* masked */, TRUE, 'O''****ly' /* masked */" in safe.debug_sql + assert safe.masked_parameters == [True, False] + assert sql_log_projection(safe) is safe + assert raw.params[0].val == "O'Reilly" + safe.params[1]._data = False + assert raw.params[1].val is True + assert 'FALSE' in sql_log_projection(safe).debug_sql + + +@pytest.mark.parametrize('value', [1, 'customer']) +def test_copied_projection_keeps_compiled_sql_structure(value): + sql = 'SELECT id FROM customer WHERE name = ? LIMIT 10000' + safe = sql_log_projection(entry(sql, [value], sql_origin='generated')) + assert safe.sql == sql + assert 'FROM customer' in sql_log_projection(replace(safe)).debug_sql + assert 'LIMIT 10000' in sql_log_projection(replace(safe)).debug_sql + + +@pytest.mark.parametrize('debug', [False, True]) +def test_inherited_intent_never_retains_raw_source(monkeypatch, debug): + from dataclasses import asdict + from teaql.sql.types import CompiledQuery + if debug: + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + source = CompiledQuery('UPDATE customer SET a=?, b=?, c=?, d=?', + ['Riverside', 'UNKNOWN-CANARY', {'api_key':'NESTED-CANARY'}, 12345.0], + parameter_log_policies=['masked','unknown','plain','masked'], sql_origin='generated') + raw = entry('SELECT id FROM customer WHERE id=?', [1], + parameter_log_policies=['plain'], sql_origin='generated') + raw.audit_reason = 'what: update Riverside UNKNOWN-CANARY NESTED-CANARY 12345.0' + raw.trace_path = [TraceNode(comment=raw.audit_reason)] + safe = sql_log_projection(raw, _intent_source=source) + assert ('Riverside' in safe.audit_reason) == debug + assert ('12345.0' in safe.audit_reason) == debug + for secret in ['UNKNOWN-CANARY','NESTED-CANARY']: + assert secret not in repr(asdict(safe)) + assert '_intent_source' not in vars(safe) + assert source.params[0] == 'Riverside' + if not debug: + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + assert 'Riverside' not in repr(sql_log_projection(replace(safe))) + assert sql_log_projection(replace(safe)).log_mode == 'masked' + + +@pytest.mark.parametrize('sql', ['SELECT $0', 'SELECT $2', 'SELECT ?, ?', 'SELECT 1', '']) +def test_invalid_bindings_are_omitted_with_reason(sql): + safe = sql_log_projection(entry(sql, ['BIND-CANARY'])) + assert safe.omission_reason is not None + assert safe.debug_sql == '[REDACTED SQL; NOT REPLAYABLE]' + assert 'BIND-CANARY' not in repr(safe) + + +def test_trusted_literals_and_backticks_ignore_false_placeholders(): + safe = sql_log_projection(entry("SELECT `field?`, 'fixed?', ? /* fixed ? */", ['Riverside'], + sql_origin='generated', parameter_log_policies=['masked'])) + assert "`field?`, 'fixed?', 'Ri*****de' /* masked */" in safe.debug_sql + + +def test_compiled_postgres_kind_and_field_policy_reach_dialect_renderer(): + dialect = PostgresDialect() + entity = (EntityDescriptor('Customer').property(PropertyDescriptor('display_name', DataType.Text)) + .audit_mask_fields(['display_name'])) + compiled = dialect.compile_select(entity, SelectQuery('Customer') + .filter(Expr.eq('display_name', 'Riverside')).limit(2)) + safe = sql_log_projection(entry(compiled.sql, compiled.params, database_kind=dialect.kind(), + parameter_log_policies=compiled.parameter_log_policies, + sql_origin=compiled.sql_origin)) + assert 'Ri*****de' in safe.debug_sql + assert 'Riverside' not in repr(safe) + assert safe.omission_reason is None + + +def test_mysql_log_renderer_uses_masks_and_skips_quoted_placeholders(): + safe = sql_log_projection(entry('SELECT `col%s`, %s, %s', ['Riverside', True], + database_kind=DatabaseKind.MySql, sql_origin='generated', + parameter_log_policies=['masked', 'plain'])) + assert "`col%s`, 'Ri*****de' /* masked */, TRUE" in safe.debug_sql + assert safe.omission_reason is None + + +def test_debug_never_reveals_inline_or_bound_credentials(monkeypatch): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + raw = entry('UPDATE users SET password = ?', ['CREDENTIAL-CANARY'], parameter_log_policies=['plain']) + assert 'CREDENTIAL-CANARY' not in repr(sql_log_projection(raw)) + raw = entry("UPDATE users SET password = 'CREDENTIAL-CANARY'", []) + safe = sql_log_projection(raw) + assert 'CREDENTIAL-CANARY' not in repr(safe) + assert safe.omission_reason == 'untrusted-literal-sql' + + +@pytest.mark.parametrize('debug', [False, True]) +def test_generated_credential_column_does_not_override_other_field_policies(monkeypatch, debug): + from teaql.provider.sqlite.dialect import SqliteDialect + if debug: + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + entity = EntityDescriptor('Customer').table_name('customer_data') + for field in ['display_name', 'public_address', 'password_hash']: + entity.property(PropertyDescriptor(field, DataType.Text).log_policy('plain')) + entity.audit_mask_fields(['display_name', 'password_hash']) + compiled = SqliteDialect().compile_select(entity, SelectQuery('Customer') + .filter(Expr.eq('display_name', 'Riverside')) + .and_filter(Expr.eq('public_address', '1 Runtime Road')) + .and_filter(Expr.eq('password_hash', 'PASSWORD-CANARY')).limit(1)) + raw = entry(compiled.sql, compiled.params, sql_origin=compiled.sql_origin, + parameter_log_policies=compiled.parameter_log_policies) + safe = sql_log_projection(raw) + assert safe.parameter_log_policies == ['masked', 'plain', 'credential'] + assert ('Riverside' if debug else 'Ri*****de') in safe.debug_sql + assert '1 Runtime Road' in safe.debug_sql + assert 'PASSWORD-CANARY' not in repr(safe) + output = [] + TextDiagnosticSqlLogSink(output.append).write(safe) + assert safe.debug_sql in output[0] + assert raw.params[2].val == 'PASSWORD-CANARY' + + +@pytest.mark.parametrize('policies', [None, ['unknown'], ['unsupported-policy']]) +def test_debug_never_exposes_unclassified_bindings(monkeypatch, policies): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + raw = entry('SELECT ?', ['UNKNOWN-BINDING-CANARY'], parameter_log_policies=policies) + raw.comment = 'what: locate UNKNOWN-BINDING-CANARY' + safe = sql_log_projection(raw) + assert 'UNKNOWN-BINDING-CANARY' not in repr(safe) + assert safe.masked_parameters == [True] + assert 'SELECT' in safe.debug_sql and 'NOT REPLAYABLE' in safe.debug_sql + assert raw.params == ['UNKNOWN-BINDING-CANARY'] + + +def test_debug_only_exposes_explicit_business_policies(monkeypatch): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + safe = sql_log_projection(entry('SELECT ?, ?, ?, ?', + ['Ordinary', 'Riverside', 'UNKNOWN-BINDING-CANARY', 'CREDENTIAL-CANARY'], + parameter_log_policies=['plain', 'masked', 'unknown', 'credential'])) + assert safe.params == ['Ordinary', 'Riverside', '[REDACTED]', '[REDACTED]'] + assert safe.masked_parameters == [False, False, True, True] + + +def test_debug_disabled_reprojects_without_recovering_plaintext(monkeypatch): + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + debug = sql_log_projection(entry('SELECT ?', ['Riverside'], parameter_log_policies=['masked'])) + assert 'Riverside' in debug.debug_sql + monkeypatch.delenv(PLAINTEXT_ENV) + safe = sql_log_projection(debug) + assert 'Riverside' not in repr(safe) + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + assert sql_log_projection(safe) is safe + + +@pytest.mark.asyncio +async def test_sqlite_crud_routes_compiler_policies_before_every_sink(tmp_path): + entity = (EntityDescriptor('Customer').table_name('customer_data') + .property(PropertyDescriptor('id', DataType.I64).is_id()) + .property(PropertyDescriptor('version', DataType.I64).is_version()) + .property(PropertyDescriptor('display_name', DataType.Text)) + .property(PropertyDescriptor('active', DataType.Bool).log_policy('plain')) + .audit_mask_fields(['display_name'])) + provider = SimpleSchemaProvider() + provider.register_entity(entity) + service = create_sqlite_service(str(tmp_path / 'mask.db'), provider) + context = RuntimeModule.new().entity(entity).into_context().with_schema_provider(service) + await context.ensure_schema() + entries, output = [], [] + sink = TextDiagnosticSqlLogSink(output.append) + def capture(log): + entries.append(log) + sink.write(log) + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=capture)) + async def mutate(command): + command.trace_chain = [TraceNode(comment='what: verify mutation log policy')] + return await service.mutate(context, MutationRequest(command)) + await mutate(InsertCommand('Customer').value('id', 1).value('version', 1) + .value('display_name', 'Riverside').value('active', True)) + query = SelectQuery('Customer').filter(Expr.new_and( + Expr.eq('display_name', 'Riverside'), Expr.eq('active', True))).limit(1) + rows = (await service.query(context, QueryRequest(query) + .comment('what: read Riverside').purpose('why: check field policy'))).rows + assert rows[0]['display_name'] == 'Riverside' + assert entries[-1].parameter_log_policies == ['masked', 'plain'] + assert entries[-1].masked_parameters == [True, False] + await mutate(UpdateCommand('Customer', Value.I64(1)).expected_version(1).value('display_name', "O'Reilly")) + await mutate(DeleteCommand('Customer', Value.I64(1)).expected_version(2)) + logs = '\n'.join(output) + assert 'Ri*****de' in logs and "O''****ly" in logs + assert 'Riverside' not in logs and "O''Reilly" not in logs + assert 'LIMIT 1' in logs and '1 rows returned' in logs + assert 'Parameterized SQL:' not in logs and 'REDACTED SQL' not in logs + assert 'Riverside' not in repr(context.sql_logs()) + assert len(entries) == 4 From e0f5bb0aa38d2655f3adfcc232c08f23889685b2 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 11:07:42 +0800 Subject: [PATCH 02/13] fix: restore MySQL dialect and verify live SQL masking --- src/teaql/provider/mysql/dialect.py | 36 ++---- .../test_sql_masking_real_databases.py | 118 ++++++++++++++++++ 2 files changed, 127 insertions(+), 27 deletions(-) create mode 100644 tests/provider/test_sql_masking_real_databases.py diff --git a/src/teaql/provider/mysql/dialect.py b/src/teaql/provider/mysql/dialect.py index 3f8eadb..6602e2b 100644 --- a/src/teaql/provider/mysql/dialect.py +++ b/src/teaql/provider/mysql/dialect.py @@ -1,31 +1,13 @@ -from teaql.sql.dialect import SqlDialect -from teaql.core.meta import EntityDescriptor +from teaql.sql.dialect import SqlDialect, quote_identifier_if_needed +from teaql.sql.types import DatabaseKind + class MysqlDialect(SqlDialect): - def quote_identifier(self, identifier: str) -> str: - return f"`{identifier}`" - + def kind(self) -> DatabaseKind: + return DatabaseKind.MySql + + def quote_ident(self, ident: str) -> str: + return quote_identifier_if_needed(ident, "`") + def placeholder(self, index: int) -> str: return "%s" - - def compile_create_table(self, entity: EntityDescriptor) -> str: - lines = [] - lines.append(f"CREATE TABLE IF NOT EXISTS {self.quote_identifier(entity.table_name)} (") - cols = [] - for prop in entity.properties: - col_type = "TEXT" - if prop.type == "U64" or prop.type == "I64": - col_type = "BIGINT" - elif prop.type == "Timestamp": - col_type = "BIGINT" - elif prop.type == "Bool": - col_type = "TINYINT(1)" - - if prop.name == "id": - cols.append(f" {self.quote_identifier(prop.column_name)} {col_type} PRIMARY KEY") - else: - cols.append(f" {self.quote_identifier(prop.column_name)} {col_type}") - - lines.append(",\n".join(cols)) - lines.append(")") - return "\n".join(lines) diff --git a/tests/provider/test_sql_masking_real_databases.py b/tests/provider/test_sql_masking_real_databases.py new file mode 100644 index 0000000..62bfd98 --- /dev/null +++ b/tests/provider/test_sql_masking_real_databases.py @@ -0,0 +1,118 @@ +"""Opt-in live PostgreSQL/MySQL masking contract for the SQL providers.""" + +import os +from types import SimpleNamespace +from uuid import uuid4 + +import pytest + +from teaql.core.expr import Expr +from teaql.core.meta import EntityDescriptor, PropertyDescriptor +from teaql.core.mutation import InsertCommand, MutationRequest, TraceNode +from teaql.core.query import SelectQuery +from teaql.core.value import DataType +from teaql.data_service import QueryRequest +from teaql.provider.mysql.dialect import MysqlDialect +from teaql.provider.mysql.transport import MysqlTransport +from teaql.provider.postgres.dialect import PostgresDialect +from teaql.provider.postgres.transport import PostgresTransport +from teaql.provider.sqlite import SimpleSchemaProvider +from teaql.runtime import RuntimeModule +from teaql.runtime.context import SqlLogOperation, TextDiagnosticSqlLogSink +from teaql.sql.executor import SqlDataServiceExecutor +from teaql.sql.types import CompiledQuery, DatabaseKind + + +def test_mysql_dialect_implements_sql_contract(): + dialect = MysqlDialect() + assert dialect.kind() == DatabaseKind.MySql + assert dialect.quote_ident("order") == "`order`" + assert dialect.placeholder(1) == "%s" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("env_name", "dialect_type", "transport_type"), + [ + ("TEAQL_TEST_POSTGRES_URL", PostgresDialect, PostgresTransport), + ("TEAQL_TEST_MYSQL_URL", MysqlDialect, MysqlTransport), + ], +) +async def test_live_provider_keeps_values_and_masks_q_and_mutation( + env_name, dialect_type, transport_type +): + url = os.getenv(env_name) + if not url: + pytest.skip(f"{env_name} is not set") + + table = f"teaql_mask_{uuid4().hex[:12]}" + entity = ( + EntityDescriptor("Customer") + .table_name(table) + .property(PropertyDescriptor("id", DataType.I64).is_id()) + .property(PropertyDescriptor("version", DataType.I64).is_version()) + .property(PropertyDescriptor("display_name", DataType.Text)) + .property(PropertyDescriptor("public_address", DataType.Text).log_policy("plain")) + .property(PropertyDescriptor("password_hash", DataType.Text)) + .audit_mask_fields(["display_name", "password_hash"]) + ) + provider = SimpleSchemaProvider() + provider.register_entity(entity) + transport = transport_type(url) + service = SqlDataServiceExecutor(dialect_type(), transport, provider) + context = RuntimeModule.new().entity(entity).into_context().with_schema_provider(service) + lines = [] + entries = [] + sink = TextDiagnosticSqlLogSink(lines.append) + + def capture(entry): + entries.append(entry) + sink.write(entry) + + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=capture)) + + try: + await context.ensure_schema() + command = ( + InsertCommand("Customer") + .value("id", 1) + .value("version", 1) + .value("display_name", "Riverside") + .value("public_address", "1 Runtime Road") + .value("password_hash", "PASSWORD-CANARY") + ) + command.trace_chain = [TraceNode(comment="what: create masked customer")] + await service.mutate(context, MutationRequest(command)) + query = SelectQuery("Customer").filter( + Expr.new_and( + Expr.eq("display_name", "Riverside"), + Expr.eq("public_address", "1 Runtime Road"), + ) + ).limit(1) + rows = ( + await service.query( + context, + QueryRequest(query) + .comment("what: read masked customer") + .purpose("why: verify live-provider SQL masking"), + ) + ).rows + assert len(rows) == 1 + assert rows[0]["display_name"] == "Riverside" + assert rows[0]["public_address"] == "1 Runtime Road" + assert [entry.operation for entry in entries[-2:]] == [ + SqlLogOperation.Insert, + SqlLogOperation.Select, + ] + assert entries[-1].parameter_log_policies == ["masked", "plain"] + assert entries[-1].masked_parameters == [True, False] + logged = "\n".join(lines) + assert "Ri*****de" in logged + assert "1 Runtime Road" in logged + assert "what: read masked customer" in logged + assert "why: verify live-provider SQL masking" in logged + assert "SELECT" in logged + assert "Riverside" not in logged + assert "PASSWORD-CANARY" not in logged + finally: + await transport.execute_sql(CompiledQuery(f"DROP TABLE IF EXISTS {table}", [])) From 9f856712d9eb5975c1e6d9b340a9b62aaa1e75bc Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 11:35:03 +0800 Subject: [PATCH 03/13] test: require live database configuration in masking gate --- tests/provider/test_sql_masking_real_databases.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/tests/provider/test_sql_masking_real_databases.py b/tests/provider/test_sql_masking_real_databases.py index 62bfd98..f36c3ef 100644 --- a/tests/provider/test_sql_masking_real_databases.py +++ b/tests/provider/test_sql_masking_real_databases.py @@ -43,6 +43,8 @@ async def test_live_provider_keeps_values_and_masks_q_and_mutation( ): url = os.getenv(env_name) if not url: + if os.getenv("TEAQL_REQUIRE_LIVE_DB", "").lower() == "true": + pytest.fail(f"{env_name} is required for live provider tests") pytest.skip(f"{env_name} is not set") table = f"teaql_mask_{uuid4().hex[:12]}" From f457ba96faecd3448ae1c64cc0f7f28664a973c2 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 14:47:36 +0800 Subject: [PATCH 04/13] fix(sql): keep schema failures out of plaintext logs --- src/teaql/sql/executor.py | 13 ++++++++---- tests/runtime/test_sql_mask_lifecycle.py | 26 ++++++++++++++++++++++++ 2 files changed, 35 insertions(+), 4 deletions(-) diff --git a/src/teaql/sql/executor.py b/src/teaql/sql/executor.py index 4e6d8fb..b6b29b7 100644 --- a/src/teaql/sql/executor.py +++ b/src/teaql/sql/executor.py @@ -4,6 +4,7 @@ from datetime import datetime import asyncio import hashlib +import logging import time import threading from array import array @@ -993,10 +994,14 @@ async def _ensure_schema(self, context: 'UserContext', capability: object) -> No await self.transport.execute_sql(CompiledQuery(idx_sql, [])) except Exception: pass - except Exception as e: - # If creating table fails, it might be due to dialect unsupported features, just pass for now - print(f"Error creating table for entity {getattr(entity, '_name', entity)}: {e}") - pass + except Exception as error: + # Preserve the existing best-effort schema behavior, but never + # forward driver exception text to an uncontrolled log sink. + logging.getLogger("teaql.sql").warning( + "Schema creation failed for entity %s (%s)", + getattr(entity, '_name', type(entity).__name__), + type(error).__name__, + ) await self.transport.execute_sql(CompiledQuery( "CREATE TABLE IF NOT EXISTS teaql_id_space (" "type_name VARCHAR(255) NOT NULL PRIMARY KEY, " diff --git a/tests/runtime/test_sql_mask_lifecycle.py b/tests/runtime/test_sql_mask_lifecycle.py index 2b83b58..924cf19 100644 --- a/tests/runtime/test_sql_mask_lifecycle.py +++ b/tests/runtime/test_sql_mask_lifecycle.py @@ -126,6 +126,32 @@ async def test_failed_statement_logs_and_preserves_error(fixture, operation): assert 'Riverside' in [p.val for p in transport.params] +@pytest.mark.asyncio +@pytest.mark.parametrize('plaintext_debug', [False, True]) +async def test_schema_failure_does_not_print_driver_error(fixture, monkeypatch, + caplog, capsys, plaintext_debug): + from teaql.runtime._schema_capability import SCHEMA_CAPABILITY + + if plaintext_debug: + monkeypatch.setenv(PLAINTEXT_ENV, PLAINTEXT_ACK) + + class SchemaFaultTransport(FaultTransport): + async def execute_sql(self, query): + if 'CREATE TABLE' in query.sql and 'teaql_id_space' not in query.sql: + raise RuntimeError('DRIVER-CANARY Riverside PASSWORD-CANARY') + return 0, None + + context, provider, _, _ = fixture + executor = SqlDataServiceExecutor(SqliteDialect(), SchemaFaultTransport(), provider) + await executor._ensure_schema(context, SCHEMA_CAPABILITY) + + printed = capsys.readouterr() + assert printed.out == printed.err == '' + assert 'Schema creation failed for entity Customer (RuntimeError)' in caplog.text + for secret in ('DRIVER-CANARY', 'Riverside', 'PASSWORD-CANARY'): + assert secret not in caplog.text + + @pytest.mark.asyncio @pytest.mark.parametrize('batches,fail,count', [(3,False,3),(0,False,0),(3,True,2),(1,True,0)]) async def test_stream_completion_and_failure(fixture, batches, fail, count): From 943d1bf1a43a4380e2ccb8df21ae2466ed578a93 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 17:27:27 +0800 Subject: [PATCH 05/13] fix(core): redact dynamic search paths in default warnings --- src/teaql/core/dynamic_search.py | 4 ++-- tests/core/test_dynamic_search.py | 11 +++++++++++ 2 files changed, 13 insertions(+), 2 deletions(-) diff --git a/src/teaql/core/dynamic_search.py b/src/teaql/core/dynamic_search.py index 0dea5b6..0d80ed8 100644 --- a/src/teaql/core/dynamic_search.py +++ b/src/teaql/core/dynamic_search.py @@ -17,8 +17,8 @@ def _reject_json_constant(_): def _warning(warning): - _LOG.warning('%s entity=%s clause=%s fieldPath=%s', warning['code'], - warning['entity'], warning['clause'], warning['fieldPath']) + _LOG.warning('%s entity=%s clause=%s fieldPath=', warning['code'], + warning['entity'], warning['clause']) def normalize_dynamic_search(source, entity, models, warn=_warning, max_clauses=100): diff --git a/tests/core/test_dynamic_search.py b/tests/core/test_dynamic_search.py index 76f6fe3..63f24db 100644 --- a/tests/core/test_dynamic_search.py +++ b/tests/core/test_dynamic_search.py @@ -25,6 +25,17 @@ def test_unknown_clauses_are_atomic_and_value_free(): assert 'SECRET' not in json.dumps(recorded) +def test_default_warning_log_omits_untrusted_path_but_keeps_structured_warning(caplog): + path = 'CLIENT_SECRET_FIELD_PATH_91' + with caplog.at_level('WARNING', logger='teaql.core.dynamic_search'): + _, warnings = normalize_dynamic_search({'filter': {path: 'SECRET_VALUE_99'}}, 'Order', MODELS) + assert warnings[0]['fieldPath'] == path + assert 'DYNAMIC_SEARCH_UNKNOWN_FIELD' in caplog.text + assert 'fieldPath=' in caplog.text + assert path not in caplog.text + assert 'SECRET_VALUE_99' not in caplog.text + + @pytest.mark.parametrize('source', ['{', '[]', 'null', '{} {}', '{"tenant":1}', '{"hardLimit":1}', '{"filter":{"removed":NaN}}', {'filter': {'name': {'$invented': 1}}}, From 87cbb33ed10ea237c480aecdc4953974cc730e42 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 20:08:55 +0800 Subject: [PATCH 06/13] fix(providers): support transactional Python graph saves --- src/teaql/provider/mysql/transport.py | 49 ++++++++++++++++++- src/teaql/provider/postgres/transport.py | 48 +++++++++++++++++- .../test_sql_masking_real_databases.py | 36 ++++++++++++++ 3 files changed, 129 insertions(+), 4 deletions(-) diff --git a/src/teaql/provider/mysql/transport.py b/src/teaql/provider/mysql/transport.py index 0b8238f..9c1f461 100644 --- a/src/teaql/provider/mysql/transport.py +++ b/src/teaql/provider/mysql/transport.py @@ -3,11 +3,11 @@ from decimal import Decimal from datetime import date, datetime, timezone from typing import List, Dict, Any, Optional, AsyncIterator -from teaql.sql.executor import SqlTransport +from teaql.sql.executor import SqlTransactionTransport, SqlTransactionTransportTx from teaql.sql.types import CompiledQuery from teaql.core.value import Value, DataType, Timestamp -class MysqlTransport(SqlTransport): +class MysqlTransport(SqlTransactionTransport): def __init__(self, db_url: str): self.db_url = db_url @@ -99,3 +99,48 @@ async def execute_sql(self, query: CompiledQuery) -> tuple[int, int]: return cur.rowcount, cur.lastrowid finally: conn.close() + + async def begin_sql(self) -> SqlTransactionTransportTx: + user, password, host, port, db = self._parse_url() + conn = await aiomysql.connect( + host=host, port=port, user=user, password=password, db=db, + cursorclass=aiomysql.DictCursor, autocommit=False, + ) + try: + await conn.begin() + return _MysqlTransaction(conn, self) + except BaseException: + conn.close() + raise + + +class _MysqlTransaction(SqlTransactionTransportTx): + def __init__(self, conn, owner: MysqlTransport): + self._conn = conn + self._owner = owner + + async def fetch_all_sql(self, query: CompiledQuery) -> List[Dict[str, Any]]: + async with self._conn.cursor() as cursor: + await cursor.execute(query.sql_with_comment(), self._owner._bind_values(query.params)) + rows = await cursor.fetchall() + return [ + {key: self._owner._decode_value(value) for key, value in row.items()} + for row in rows + ] + + async def execute_sql(self, query: CompiledQuery) -> tuple[int, int]: + async with self._conn.cursor() as cursor: + await cursor.execute(query.sql_with_comment(), self._owner._bind_values(query.params)) + return cursor.rowcount, cursor.lastrowid + + async def commit_sql(self) -> None: + try: + await self._conn.commit() + finally: + self._conn.close() + + async def rollback_sql(self) -> None: + try: + await self._conn.rollback() + finally: + self._conn.close() diff --git a/src/teaql/provider/postgres/transport.py b/src/teaql/provider/postgres/transport.py index 7627b51..f8e373f 100644 --- a/src/teaql/provider/postgres/transport.py +++ b/src/teaql/provider/postgres/transport.py @@ -3,11 +3,11 @@ from decimal import Decimal from datetime import date, datetime, timezone from typing import List, Dict, Any, Optional, AsyncIterator -from teaql.sql.executor import SqlTransport +from teaql.sql.executor import SqlTransactionTransport, SqlTransactionTransportTx from teaql.sql.types import CompiledQuery from teaql.core.value import Value, DataType, Timestamp -class PostgresTransport(SqlTransport): +class PostgresTransport(SqlTransactionTransport): def __init__(self, db_url: str): self.db_url = db_url @@ -88,3 +88,47 @@ async def execute_sql(self, query: CompiledQuery) -> tuple[int, int]: return affected_rows, 0 finally: await conn.close() + + async def begin_sql(self) -> SqlTransactionTransportTx: + conn = await asyncpg.connect(self.db_url) + try: + transaction = conn.transaction() + await transaction.start() + return _PostgresTransaction(conn, transaction, self) + except BaseException: + await conn.close() + raise + + +class _PostgresTransaction(SqlTransactionTransportTx): + def __init__(self, conn, transaction, owner: PostgresTransport): + self._conn = conn + self._transaction = transaction + self._owner = owner + + async def fetch_all_sql(self, query: CompiledQuery) -> List[Dict[str, Any]]: + rows = await self._conn.fetch(query.sql_with_comment(), *self._owner._bind_values(query.params)) + return [ + {key: self._owner._decode_value(value) for key, value in row.items()} + for row in rows + ] + + async def execute_sql(self, query: CompiledQuery) -> tuple[int, int]: + status = await self._conn.execute(query.sql_with_comment(), *self._owner._bind_values(query.params)) + if status.startswith("INSERT "): + return int(status.split()[-1]), 0 + if status.startswith(("UPDATE ", "DELETE ")): + return int(status.split()[-1]), 0 + return 0, 0 + + async def commit_sql(self) -> None: + try: + await self._transaction.commit() + finally: + await self._conn.close() + + async def rollback_sql(self) -> None: + try: + await self._transaction.rollback() + finally: + await self._conn.close() diff --git a/tests/provider/test_sql_masking_real_databases.py b/tests/provider/test_sql_masking_real_databases.py index f36c3ef..daed979 100644 --- a/tests/provider/test_sql_masking_real_databases.py +++ b/tests/provider/test_sql_masking_real_databases.py @@ -118,3 +118,39 @@ def capture(entry): assert "PASSWORD-CANARY" not in logged finally: await transport.execute_sql(CompiledQuery(f"DROP TABLE IF EXISTS {table}", [])) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("env_name", "transport_type"), + [ + ("TEAQL_TEST_POSTGRES_URL", PostgresTransport), + ("TEAQL_TEST_MYSQL_URL", MysqlTransport), + ], +) +async def test_live_provider_transaction_commit_and_rollback(env_name, transport_type): + url = os.getenv(env_name) + if not url: + if os.getenv("TEAQL_REQUIRE_LIVE_DB", "").lower() == "true": + pytest.fail(f"{env_name} is required for live provider tests") + pytest.skip(f"{env_name} is not set") + + table = f"teaql_tx_{uuid4().hex[:12]}" + transport = transport_type(url) + await transport.execute_sql(CompiledQuery( + f"CREATE TABLE {table} (id BIGINT PRIMARY KEY, display_name VARCHAR(100))", [])) + try: + transaction = await transport.begin_sql() + await transaction.execute_sql(CompiledQuery( + f"INSERT INTO {table} (id, display_name) VALUES (1, 'rolled back')", [])) + assert len(await transaction.fetch_all_sql(CompiledQuery(f"SELECT id FROM {table}", []))) == 1 + await transaction.rollback_sql() + assert await transport.fetch_all_sql(CompiledQuery(f"SELECT id FROM {table}", [])) == [] + + transaction = await transport.begin_sql() + await transaction.execute_sql(CompiledQuery( + f"INSERT INTO {table} (id, display_name) VALUES (2, 'committed')", [])) + await transaction.commit_sql() + assert len(await transport.fetch_all_sql(CompiledQuery(f"SELECT id FROM {table}", []))) == 1 + finally: + await transport.execute_sql(CompiledQuery(f"DROP TABLE IF EXISTS {table}", [])) From c40d938656ca15e5cb2f871704bc6422121f239a Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 21:21:05 +0800 Subject: [PATCH 07/13] fix: inherit query intent in legacy generated Python requests --- src/teaql/data_service/__init__.py | 9 +++++++++ tests/runtime/test_sql_mask_lifecycle.py | 20 ++++++++++++++++++++ 2 files changed, 29 insertions(+) diff --git a/src/teaql/data_service/__init__.py b/src/teaql/data_service/__init__.py index 5bcbe59..bb1c255 100644 --- a/src/teaql/data_service/__init__.py +++ b/src/teaql/data_service/__init__.py @@ -32,6 +32,15 @@ class QueryRequest: _comment: Optional[str] = None _purpose: Optional[str] = None + def __post_init__(self) -> None: + # Older generated wrappers store intent on SelectQuery but construct + # QueryRequest(query) without forwarding it. Keep the explicit request + # values authoritative while preserving their diagnostic intent. + if self._comment is None: + self._comment = getattr(self.query, 'comment_text', None) + if self._purpose is None: + self._purpose = getattr(self.query, 'purpose_text', None) + def comment(self, text: str) -> 'QueryRequest': if not text: raise ValueError("comment cannot be empty") diff --git a/tests/runtime/test_sql_mask_lifecycle.py b/tests/runtime/test_sql_mask_lifecycle.py index 924cf19..04e6da8 100644 --- a/tests/runtime/test_sql_mask_lifecycle.py +++ b/tests/runtime/test_sql_mask_lifecycle.py @@ -84,6 +84,26 @@ def request(): return QueryRequest(query).comment('what: inspect customers').purpose('why: lifecycle test') +@pytest.mark.asyncio +async def test_old_generated_query_request_keeps_query_intent(fixture): + context, provider, output, entries = fixture + query = (SelectQuery('Customer').filter(Expr.eq('display_name', 'Riverside')) + .limit(5).comment('what: old generated request') + .purpose('why: preserve diagnostic intent')) + request = QueryRequest(query) + assert request._comment == 'what: old generated request' + assert request._purpose == 'why: preserve diagnostic intent' + assert QueryRequest(query, _comment='explicit what', _purpose='explicit why')._comment == 'explicit what' + assert QueryRequest(query, _comment='explicit what', _purpose='explicit why')._purpose == 'explicit why' + + executor = DomainStatementExecutor(SqliteDialect(), FaultTransport(), provider) + await executor.query(context, request) + assert entries[0].comment == 'what: old generated request' + assert entries[0].purpose == 'why: preserve diagnostic intent' + assert 'what: old generated request' in output[0] + assert 'why: preserve diagnostic intent' in output[0] + + def assert_log(fixture, outcome, count=None, debug=False): context, _, output, entries = fixture assert len(entries) == 1 From 735f8002107ad9cf27bf3777bb11f99e8d9e50fc Mon Sep 17 00:00:00 2001 From: Philip Z Date: Mon, 28 Sep 2026 23:15:41 +0800 Subject: [PATCH 08/13] Fail closed on missing Python generated SQL log metadata --- src/teaql/core/meta.py | 4 +++- src/teaql/sql/dialect.py | 2 ++ tests/runtime/test_sql_masking_policy.py | 18 ++++++++++++++++++ 3 files changed, 23 insertions(+), 1 deletion(-) diff --git a/src/teaql/core/meta.py b/src/teaql/core/meta.py index 18cedca..50c9f3d 100644 --- a/src/teaql/core/meta.py +++ b/src/teaql/core/meta.py @@ -53,6 +53,7 @@ def __init__(self, name: str): self.properties = [] self.relations = [] self.audit_mask_fields_val = [] + self.audit_mask_fields_declared = False self.audit_value_max_len_val = None def table_name(self, name): self.table_name_val = name @@ -72,7 +73,8 @@ def property_by_name(self, name): return next((prop for prop in self.properties if prop.name == name), None) def audit_mask_fields(self, fields): - self.audit_mask_fields_val = list(fields) + self.audit_mask_fields_val = list(fields) if fields is not None else [] + self.audit_mask_fields_declared = fields is not None return self def audit_value_max_len(self, max_len): diff --git a/src/teaql/sql/dialect.py b/src/teaql/sql/dialect.py index a7827c1..b29d901 100644 --- a/src/teaql/sql/dialect.py +++ b/src/teaql/sql/dialect.py @@ -403,6 +403,8 @@ def field_log_policy(self, entity, field): prop = entity.property_by_name(field) if credential_name(field) or (prop and credential_name(prop.column_name_val)): return 'credential' + if not entity.audit_mask_fields_declared: + return 'unknown' if field in entity.audit_mask_fields_val: return 'masked' return getattr(prop, 'log_policy_val', 'unknown') diff --git a/tests/runtime/test_sql_masking_policy.py b/tests/runtime/test_sql_masking_policy.py index 9c3ca75..8dc0462 100644 --- a/tests/runtime/test_sql_masking_policy.py +++ b/tests/runtime/test_sql_masking_policy.py @@ -105,6 +105,24 @@ def test_compiled_postgres_kind_and_field_policy_reach_dialect_renderer(): assert safe.omission_reason is None +def test_legacy_missing_mask_metadata_overrides_old_plain_property_policy(): + from teaql.provider.sqlite.dialect import SqliteDialect + entity = EntityDescriptor('Customer').table_name('customer_data') + entity.property(PropertyDescriptor('name', DataType.Text).log_policy('plain')) + query = SelectQuery('Customer').filter(Expr.eq('name', 'PRIVATE-CANARY')).limit(1) + dialect = SqliteDialect() + compiled = dialect.compile_select(entity, query) + assert compiled.parameter_log_policies == ['unknown'] + safe = sql_log_projection(entry(compiled.sql, compiled.params, + sql_origin=compiled.sql_origin, + parameter_log_policies=compiled.parameter_log_policies)) + assert 'PRIVATE-CANARY' not in repr(safe) + assert '[REDACTED]' in safe.debug_sql + entity.audit_mask_fields([]) + declared = dialect.compile_select(entity, query) + assert declared.parameter_log_policies == ['plain'] + + def test_mysql_log_renderer_uses_masks_and_skips_quoted_placeholders(): safe = sql_log_projection(entry('SELECT `col%s`, %s, %s', ['Riverside', True], database_kind=DatabaseKind.MySql, sql_origin='generated', From 57874a8db1802aa37361438be0202992acad84e0 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Tue, 29 Sep 2026 00:37:34 +0800 Subject: [PATCH 09/13] test: guard telemetry failure exports against driver messages --- tests/runtime/test_opentelemetry.py | 56 ++++++++++++++++++++++++++++- 1 file changed, 55 insertions(+), 1 deletion(-) diff --git a/tests/runtime/test_opentelemetry.py b/tests/runtime/test_opentelemetry.py index b42ac54..01505ec 100644 --- a/tests/runtime/test_opentelemetry.py +++ b/tests/runtime/test_opentelemetry.py @@ -1,3 +1,5 @@ +import pytest + from opentelemetry.sdk.metrics import MeterProvider from opentelemetry.sdk.metrics.export import InMemoryMetricReader from opentelemetry.sdk.trace import TracerProvider @@ -5,7 +7,9 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from teaql.runtime.opentelemetry import OpenTelemetryRuntimeTelemetry -from teaql.runtime.telemetry import RuntimeOperation, start_runtime_operation +from teaql.runtime.telemetry import ( + RuntimeOperation, observe_runtime_operation_sync, start_runtime_operation, +) def test_exports_safe_spans_and_metrics_through_official_sdk(): @@ -92,6 +96,56 @@ def test_delegates_explicit_application_owned_lifecycle(): assert calls == ["flush", "shutdown"] tracer_provider.shutdown() meter_provider.shutdown() + + +def test_failure_telemetry_does_not_export_driver_error_message(): + class DriverCanaryError(Exception): + def __str__(self): + raise AssertionError("telemetry must not format the driver exception") + + span_exporter = InMemorySpanExporter() + tracer_provider = TracerProvider() + tracer_provider.add_span_processor(SimpleSpanProcessor(span_exporter)) + meter_provider = MeterProvider() + log_exporter = InMemoryLogRecordExporter() + logger_provider = LoggerProvider() + logger_provider.add_log_record_processor(SimpleLogRecordProcessor(log_exporter)) + handler = LoggingHandler(logger_provider=logger_provider) + runtime_logger = logging.getLogger("teaql.runtime") + original_level = runtime_logger.level + runtime_logger.setLevel(logging.INFO) + runtime_logger.addHandler(handler) + telemetry = OpenTelemetryRuntimeTelemetry( + tracer_provider.get_tracer("io.teaql.runtime"), + meter_provider.get_meter("io.teaql.runtime"), + ) + error = DriverCanaryError("SQL failed for password=OTEL-FAILURE-CANARY") + + try: + with pytest.raises(DriverCanaryError) as caught: + observe_runtime_operation_sync( + telemetry, RuntimeOperation("provider", "sqlite.query"), + lambda: _raise(error), + ) + assert caught.value is error + span = span_exporter.get_finished_spans()[0] + log = log_exporter.get_finished_logs()[0].log_record + assert span.attributes["teaql.error.type"] == "DriverCanaryError" + assert log.attributes["teaql.operation.outcome"] == "failure" + exported = repr((span.attributes, span.status, span.events, + log.body, log.attributes)) + assert "OTEL-FAILURE-CANARY" not in exported + assert "password=" not in exported + finally: + runtime_logger.removeHandler(handler) + runtime_logger.setLevel(original_level) + logger_provider.shutdown() + tracer_provider.shutdown() + meter_provider.shutdown() + + +def _raise(error): + raise error import logging from opentelemetry.instrumentation.logging.handler import LoggingHandler From 1bde4398b201aafc310b2354619eb66d4ab894a0 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Tue, 29 Sep 2026 00:56:44 +0800 Subject: [PATCH 10/13] Keep SQL diagnostic sink failures out of queries --- src/teaql/runtime/context.py | 19 +++++++++++++------ tests/runtime/test_sql_mask_lifecycle.py | 19 +++++++++++++++++++ 2 files changed, 32 insertions(+), 6 deletions(-) diff --git a/src/teaql/runtime/context.py b/src/teaql/runtime/context.py index fa928fd..b1bc376 100644 --- a/src/teaql/runtime/context.py +++ b/src/teaql/runtime/context.py @@ -752,9 +752,7 @@ def _record_metadata_log(self, metadata: Any, *, intent_source=None): logs = self.sql_logs() logs.append(entry) self._resources["sql_logs"] = logs - sink = self.get_resource("diagnostic_sql_log_sink") - if sink is not None: - sink.write(entry) + self._write_diagnostic_sql_log(entry) buf = self.get_resource("UnifiedLogBuffer") if buf: buf.entries.append(UnifiedLogEntry( @@ -803,9 +801,7 @@ def record_sql_log(self, operation: Any, query: Any, started_at: Any, ended_at: logs = self.sql_logs() logs.append(entry) self._resources["sql_logs"] = logs - sink = self.get_resource("diagnostic_sql_log_sink") - if sink is not None: - sink.write(entry) + self._write_diagnostic_sql_log(entry) buf = self.get_resource("UnifiedLogBuffer") if buf: @@ -816,6 +812,17 @@ def record_sql_log(self, operation: Any, query: Any, started_at: Any, ended_at: payload=LogPayload.Sql(entry) )) + def _write_diagnostic_sql_log(self, entry: 'SqlLogEntry'): + sink = self.get_resource("diagnostic_sql_log_sink") + if sink is None: + return + try: + sink.write(entry) + except Exception: + # An optional diagnostic destination must not change SQL outcomes. + # Do not print the exception: custom sinks may include raw values. + pass + def register_executor(self, executor: Any): self.insert_resource("executor", executor) diff --git a/tests/runtime/test_sql_mask_lifecycle.py b/tests/runtime/test_sql_mask_lifecycle.py index 04e6da8..3a5b7b5 100644 --- a/tests/runtime/test_sql_mask_lifecycle.py +++ b/tests/runtime/test_sql_mask_lifecycle.py @@ -262,6 +262,25 @@ def broken(entry): assert len(context.sql_logs()) == 1 +@pytest.mark.asyncio +async def test_broken_sink_cannot_fail_successful_query(fixture): + context, provider, _, _ = fixture + attempts = [] + + def broken(entry): + attempts.append(entry) + raise RuntimeError('DIAGNOSTIC-SINK-CANARY') + + context.set_diagnostic_sql_log_sink(SimpleNamespace(write=broken)) + executor = DomainStatementExecutor(SqliteDialect(), FaultTransport(), provider) + result = await executor.query(context, request()) + + assert result is not None + assert len(attempts) == 1 + assert len(context.sql_logs()) == 1 + assert 'PASSWORD-CANARY' not in str(context.sql_logs()[0]) + + @pytest.mark.asyncio async def test_query_cancellation_preserved(fixture): context, provider, _, _ = fixture From ce36e944957abceb514d2ac6f018fa837e9771e8 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Tue, 29 Sep 2026 04:40:00 +0800 Subject: [PATCH 11/13] fix: scrub target IDs from Python app audit traces --- src/teaql/runtime/audit.py | 5 +++-- tests/runtime/test_runtime.py | 14 ++++++++++++++ 2 files changed, 17 insertions(+), 2 deletions(-) diff --git a/src/teaql/runtime/audit.py b/src/teaql/runtime/audit.py index c52cbd5..0d8096a 100644 --- a/src/teaql/runtime/audit.py +++ b/src/teaql/runtime/audit.py @@ -48,9 +48,10 @@ def safe(self, mask_fields: List[str], max_length: Optional[int]) -> "SafeAuditE if truncated: value = "*" * max_length if max_length <= 3 else value[:max_length - 3] + "..." fields.append(SafeAuditField(change.field, value, masked, truncated)) + intent_values = secrets + value_strings(self.entity_id) return SafeAuditEvent( - self.kind, self.entity, self.entity_id, scrub(tuple(fields), secrets), scrub(self.trace_chain, secrets), - scrub(self.actor, secrets), self.category, + self.kind, self.entity, self.entity_id, scrub(tuple(fields), secrets), scrub(self.trace_chain, intent_values), + scrub(self.actor, intent_values), self.category, ) diff --git a/tests/runtime/test_runtime.py b/tests/runtime/test_runtime.py index abc9c04..ede6d1f 100644 --- a/tests/runtime/test_runtime.py +++ b/tests/runtime/test_runtime.py @@ -46,6 +46,20 @@ def test_bootstrap_audit_identity_survives_safe_projection(): assert safe.category == "runtime-bootstrap" +def test_app_audit_trace_redacts_target_id_without_changing_structured_id(): + event = RawAuditEvent( + MutationAuditKind.UPDATED, + "SchoolType", + 1001, + (AuditFieldChange("name", "Primary", "Primary School"),), + ("rename SchoolType 1001",), + ) + safe = event.safe([], None) + assert safe.entity_id == 1001 + assert safe.trace_chain == ("rename SchoolType [REDACTED]",) + assert event.trace_chain == ("rename SchoolType 1001",) + + class RecordingTransaction: def __init__(self, events): self.events = events From 191f50a9876dd713f4e0acee44d9b4c429ddf997 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Tue, 29 Sep 2026 04:46:44 +0800 Subject: [PATCH 12/13] test: cover Python audit ID redaction across mutation kinds --- tests/runtime/test_runtime.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/runtime/test_runtime.py b/tests/runtime/test_runtime.py index ede6d1f..6c92975 100644 --- a/tests/runtime/test_runtime.py +++ b/tests/runtime/test_runtime.py @@ -46,9 +46,10 @@ def test_bootstrap_audit_identity_survives_safe_projection(): assert safe.category == "runtime-bootstrap" -def test_app_audit_trace_redacts_target_id_without_changing_structured_id(): +@pytest.mark.parametrize("kind", list(MutationAuditKind)) +def test_app_audit_trace_redacts_target_id_without_changing_structured_id(kind): event = RawAuditEvent( - MutationAuditKind.UPDATED, + kind, "SchoolType", 1001, (AuditFieldChange("name", "Primary", "Primary School"),), From 52499addf89bac22204df67640d9ddd7e52d85e2 Mon Sep 17 00:00:00 2001 From: Philip Z Date: Tue, 29 Sep 2026 05:07:26 +0800 Subject: [PATCH 13/13] Scrub mutation target IDs from Python SQL intent logs --- src/teaql/runtime/context.py | 4 +- src/teaql/runtime/log_privacy.py | 13 +++--- src/teaql/sql/executor.py | 17 ++++++-- tests/runtime/test_sql_mask_lifecycle.py | 50 +++++++++++++++++++++++- 4 files changed, 71 insertions(+), 13 deletions(-) diff --git a/src/teaql/runtime/context.py b/src/teaql/runtime/context.py index b1bc376..5b299d4 100644 --- a/src/teaql/runtime/context.py +++ b/src/teaql/runtime/context.py @@ -707,7 +707,7 @@ def language(self) -> Any: def record_metadata_log(self, metadata: Any): self._record_metadata_log(metadata) - def _record_metadata_log(self, metadata: Any, *, intent_source=None): + def _record_metadata_log(self, metadata: Any, *, intent_source=None, intent_values=()): """Internal statement plumbing: source bindings never reach sinks/buffers.""" op = SqlLogOperation.Select op_str = str(getattr(metadata, 'operation', '')).lower() @@ -748,7 +748,7 @@ def _record_metadata_log(self, metadata: Any, *, intent_source=None): entry.result_summary = f"{entry.affected_rows} rows affected" from .log_privacy import sql_log_projection - entry = sql_log_projection(entry, _intent_source=intent_source) + entry = sql_log_projection(entry, _intent_source=intent_source, _intent_values=intent_values) logs = self.sql_logs() logs.append(entry) self._resources["sql_logs"] = logs diff --git a/src/teaql/runtime/log_privacy.py b/src/teaql/runtime/log_privacy.py index 1f71527..eb02fdb 100644 --- a/src/teaql/runtime/log_privacy.py +++ b/src/teaql/runtime/log_privacy.py @@ -117,7 +117,7 @@ def _is_masked(policy, allow): return policy in ('credential', 'unknown') or (not allow and policy != 'plain') -def sql_log_projection(entry, *, _intent_source=None): +def sql_log_projection(entry, *, _intent_source=None, _intent_values=()): """Source bindings are call-local runtime plumbing, never stored on a log entry.""" allow = plaintext_enabled() and entry.log_mode != 'masked' prior = _projections.get(id(entry)) @@ -129,12 +129,12 @@ def sql_log_projection(entry, *, _intent_source=None): # Entries are mutable. Do not expose the cached safe alternative itself. if prior[3] is not None: return _remember_projection(deepcopy(prior[3]), False) - projected = _project_with_policy(entry, allow, _intent_source) - alternative = _project_with_policy(entry, False, _intent_source) if allow else None + projected = _project_with_policy(entry, allow, _intent_source, _intent_values) + alternative = _project_with_policy(entry, False, _intent_source, _intent_values) if allow else None return _remember_projection(projected, allow, alternative) -def _project_with_policy(entry, allow, intent_source): +def _project_with_policy(entry, allow, intent_source, intent_values): from teaql.sql.types import DatabaseKind, render_log_sql, _sql_literal from teaql.core.value import Value supplied = entry.parameter_log_policies @@ -154,11 +154,12 @@ def _project_with_policy(entry, allow, intent_source): source_policies = _binding_policies(intent_source) secrets.extend(text for index, value in enumerate(intent_source.params) if _is_masked(source_policies[index], allow) for text in value_strings(value)) + intent_secrets = secrets + [text for value in intent_values for text in value_strings(value)] # A copied/changed debug record has lost its reliable private alternative. # Its inherited intent may mention values absent from its own SQL bindings. unknown_debug_intent = not allow and entry.log_mode == 'debug-plaintext' and intent_source is None def intent(value): - return scrub(value, secrets, hide_all=unknown_debug_intent) + return scrub(value, intent_secrets, hide_all=unknown_debug_intent) bare = re.sub(r'\$[0-9]+', '?', entry.sql) unsafe = ((not allow or credentials) and entry.sql_origin != 'generated' and bool(re.search(r"['\"`$]|--|/\*|\b\d+\b|:[A-Za-z_]", bare))) @@ -189,6 +190,6 @@ def literal(index): audit_reason=intent(entry.audit_reason), result_summary=(f'{entry.result_count} rows returned' if entry.result_count is not None else f'{entry.affected_rows} rows affected' if entry.affected_rows is not None - else intent(entry.result_summary)), + else scrub(entry.result_summary, secrets, hide_all=unknown_debug_intent)), trace_path=intent(entry.trace_path)) return projected diff --git a/src/teaql/sql/executor.py b/src/teaql/sql/executor.py index b6b29b7..a8450b1 100644 --- a/src/teaql/sql/executor.py +++ b/src/teaql/sql/executor.py @@ -211,10 +211,20 @@ def _record_statement(self, context, request, compiled, started_at, operation, if context is not None: try: source = getattr(request, '_log_intent_source', None) if query else None - if source is None: + target_id = None + if not query: + target_id = getattr(request._data, 'id', None) + if target_id is None: + descriptor = self.schema_provider.get_entity(entity) + id_property = next((prop for prop in descriptor.properties + if getattr(prop, '_is_id', False) or getattr(prop, 'is_id_val', False)), None) + if id_property is not None: + target_id = getattr(request._data, 'values', {}).get(id_property.name) + if source is None and target_id is None: context.record_metadata_log(metadata) else: - context._record_metadata_log(metadata, intent_source=source) + context._record_metadata_log(metadata, intent_source=source, + intent_values=() if target_id is None else (target_id,)) except Exception: # A broken diagnostic destination must not replace an in-flight # driver failure, cancellation or generator close. @@ -867,7 +877,8 @@ def _record_readback(self, context, readback, source, write_metadata, started_at affected_rows=None, result_count=len(rows) if rows is not None else None, trace_chain=[*write_metadata.trace_chain, TraceNode(kind='sql', name='readback', comment='readback')]) try: - context._record_metadata_log(metadata, intent_source=source) + context._record_metadata_log(metadata, intent_source=source, + intent_values=tuple(readback.params[:1])) except BaseException: # An in-flight readback error must survive a diagnostic sink failure. pass diff --git a/tests/runtime/test_sql_mask_lifecycle.py b/tests/runtime/test_sql_mask_lifecycle.py index 3a5b7b5..300d11b 100644 --- a/tests/runtime/test_sql_mask_lifecycle.py +++ b/tests/runtime/test_sql_mask_lifecycle.py @@ -1,14 +1,15 @@ import asyncio +from datetime import datetime from types import SimpleNamespace import pytest from teaql.core.expr import Expr from teaql.core.meta import EntityDescriptor, PropertyDescriptor -from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, MutationRequest, TraceNode +from teaql.core.mutation import InsertCommand, UpdateCommand, DeleteCommand, RecoverCommand, MutationRequest, TraceNode from teaql.core.query import SelectQuery from teaql.core.value import DataType, Value -from teaql.data_service import QueryRequest +from teaql.data_service import DataServiceOperation, ExecutionMetadata, QueryRequest from teaql.provider.sqlite import SimpleSchemaProvider from teaql.provider.sqlite.dialect import SqliteDialect from teaql.runtime import RuntimeModule @@ -343,6 +344,51 @@ async def test_generated_mutation_string_comment_is_preserved(fixture): assert_log(fixture, 'success') +@pytest.mark.asyncio +@pytest.mark.parametrize('operation', ['insert', 'update', 'delete', 'recover']) +@pytest.mark.parametrize('failure', [False, True]) +async def test_sql_intent_scrubs_target_id_without_changing_bindings(fixture, operation, failure): + context, provider, output, entries = fixture + next(p for p in provider.get_entity('Customer').properties if p.name == 'id').log_policy('plain') + transport = FaultTransport(failure=RuntimeError('DRIVER-CANARY') if failure else None) + executor = DomainStatementExecutor(SqliteDialect(), transport, provider) + commands = { + 'insert': InsertCommand('Customer').value('id', 1001).value('display_name', 'Riverside'), + 'update': UpdateCommand('Customer', Value.I64(1001)).expected_version(1).value('display_name', 'Riverside'), + 'delete': DeleteCommand('Customer', Value.I64(1001)).expected_version(1), + 'recover': RecoverCommand('Customer', Value.I64(1001), -2), + } + command = commands[operation] + command.trace_chain = [TraceNode(comment='what: mutate target 1001')] + mutation = MutationRequest(command) + mutation.comment = 'what: mutate target 1001' + if failure: + with pytest.raises(TransportError): + await executor.mutate(context, mutation) + else: + await executor.mutate(context, mutation) + assert entries[0].audit_reason == 'what: mutate target [REDACTED]' + assert '1001' not in repr(entries[0].trace_path) + assert '1001' not in output[0].split('auditReason=', 1)[1].split('Debug SQL:', 1)[0] + assert any(getattr(value, 'val', value) == 1001 for value in transport.params) + assert mutation.comment == 'what: mutate target 1001' + + +def test_short_target_id_does_not_redact_structural_row_count(fixture): + context, _, _, entries = fixture + now = datetime.now() + metadata = ExecutionMetadata(backend='sqlite', operation=DataServiceOperation.Update, + started_at=now, ended_at=now, parameterized_sql='UPDATE customer_data SET version = ? WHERE id = ?', + parameters=[Value.I64(2), Value.I64(1)], affected_rows=1, + audit_reason='what: update target 1', database_kind=SqliteDialect().kind(), + parameter_log_policies=['plain', 'plain'], sql_origin='generated') + context._record_metadata_log(metadata, intent_values=(Value.I64(1),)) + assert entries[0].audit_reason == 'what: update target [REDACTED]' + assert entries[0].result_summary == '1 rows affected' + assert entries[0].params[1].val == 1 + assert metadata.audit_reason == 'what: update target 1' + + class ReadbackTransaction(FaultTransport, SqlTransactionTransportTx): def __init__(self, rows=None, failure=None): super().__init__(failure=failure)