diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 31cfda207b..4f963b1387 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -22,7 +22,7 @@ dependencies = [ "psycopg2-binary==2.9.11", "Faker==40.1.0", "gunicorn==23.0.0", - "uvicorn[standard]==0.40.0", + "uvicorn[standard]==0.52.4", "websockets==16.0.0", "requests==2.34.2", "itsdangerous==2.2.0", diff --git a/backend/uv.lock b/backend/uv.lock index d28a094ea3..971254ad1a 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -450,7 +450,7 @@ requires-dist = [ { name = "typing-extensions", specifier = ">=4.14.1" }, { name = "tzdata", specifier = "==2026.2" }, { name = "unicodecsv", specifier = "==0.14.1" }, - { name = "uvicorn", extras = ["standard"], specifier = "==0.40.0" }, + { name = "uvicorn", extras = ["standard"], specifier = "==0.52.4" }, { name = "validators", specifier = "==0.35.0" }, { name = "websockets", specifier = "==16.0.0" }, { name = "xlrd", specifier = "==2.0.2" }, @@ -1629,16 +1629,22 @@ wheels = [ [[package]] name = "httptools" -version = "0.7.1" +version = "0.8.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/b5/46/120a669232c7bdedb9d52d4aeae7e6c7dfe151e99dc70802e2fc7a5e1993/httptools-0.7.1.tar.gz", hash = "sha256:abd72556974f8e7c74a259655924a717a2365b236c882c3f6f8a45fe94703ac9", size = 258961, upload-time = "2025-10-10T03:55:08.559Z" } +sdist = { url = "https://files.pythonhosted.org/packages/43/e5/d471fcb0e14523fe1c3f4ba58ca52480e7bd70ad7109a3846bc75892f7fb/httptools-0.8.0.tar.gz", hash = "sha256:6b2a32f18d97e16e90827d7a819ffa8dbd8cc245fc4e1fa9d1095b54ef4bd999", size = 271342, upload-time = "2026-05-25T22:17:48.841Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/34/50/9d095fcbb6de2d523e027a2f304d4551855c2f46e0b82befd718b8b20056/httptools-0.7.1-cp314-cp314-macosx_10_13_universal2.whl", hash = "sha256:c08fe65728b8d70b6923ce31e3956f859d5e1e8548e6f22ec520a962c6757270", size = 203619, upload-time = "2025-10-10T03:54:54.321Z" }, - { url = "https://files.pythonhosted.org/packages/07/f0/89720dc5139ae54b03f861b5e2c55a37dba9a5da7d51e1e824a1f343627f/httptools-0.7.1-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:7aea2e3c3953521c3c51106ee11487a910d45586e351202474d45472db7d72d3", size = 108714, upload-time = "2025-10-10T03:54:55.163Z" }, - { url = "https://files.pythonhosted.org/packages/b3/cb/eea88506f191fb552c11787c23f9a405f4c7b0c5799bf73f2249cd4f5228/httptools-0.7.1-cp314-cp314-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:0e68b8582f4ea9166be62926077a3334064d422cf08ab87d8b74664f8e9058e1", size = 472909, upload-time = "2025-10-10T03:54:56.056Z" }, - { url = "https://files.pythonhosted.org/packages/e0/4a/a548bdfae6369c0d078bab5769f7b66f17f1bfaa6fa28f81d6be6959066b/httptools-0.7.1-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:df091cf961a3be783d6aebae963cc9b71e00d57fa6f149025075217bc6a55a7b", size = 470831, upload-time = "2025-10-10T03:54:57.219Z" }, - { url = "https://files.pythonhosted.org/packages/4d/31/14df99e1c43bd132eec921c2e7e11cda7852f65619bc0fc5bdc2d0cb126c/httptools-0.7.1-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:f084813239e1eb403ddacd06a30de3d3e09a9b76e7894dcda2b22f8a726e9c60", size = 452631, upload-time = "2025-10-10T03:54:58.219Z" }, - { url = "https://files.pythonhosted.org/packages/22/d2/b7e131f7be8d854d48cb6d048113c30f9a46dca0c9a8b08fcb3fcd588cdc/httptools-0.7.1-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:7347714368fb2b335e9063bc2b96f2f87a9ceffcd9758ac295f8bbcd3ffbc0ca", size = 452910, upload-time = "2025-10-10T03:54:59.366Z" }, + { url = "https://files.pythonhosted.org/packages/1a/12/fa3fbf5f9517b273edea2dc982aa82a8c634091e67c590792b729017bc6f/httptools-0.8.0-cp314-cp314-macosx_10_13_universal2.whl", hash = "sha256:de242a49b5d18e0a8776e654e9f6bf6d89f3875a5c35b425a0e7ce940feb3fd6", size = 206183, upload-time = "2026-05-25T22:17:24.004Z" }, + { url = "https://files.pythonhosted.org/packages/30/fc/5e7c4cb443370f2090a3aba0453a07384d29ff66b7435bb90e77e1037599/httptools-0.8.0-cp314-cp314-macosx_11_0_arm64.whl", hash = "sha256:159e9ab5f701ccd42e555a12f1ad8ff69702910fc1c996cf2bb66e5fcb7a231b", size = 112079, upload-time = "2026-05-25T22:17:25.216Z" }, + { url = "https://files.pythonhosted.org/packages/ba/53/771bd891eb0f236f32145d6a1775777ec85745f3cc983a1f23d1a3b8ddfe/httptools-0.8.0-cp314-cp314-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:c4a9f1707e4823d54dfec6c33fa3697d302aed536ed352a7ebb5a061ddb869d0", size = 481596, upload-time = "2026-05-25T22:17:26.186Z" }, + { url = "https://files.pythonhosted.org/packages/62/42/94e15bc68ce3d423243c45d7f1b0c7561f13844f97dc52ae23182fb65628/httptools-0.8.0-cp314-cp314-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:d76ad7b951387e3632c8716a9bb03ac5b45c5f16119aa409db0459520887944e", size = 480865, upload-time = "2026-05-25T22:17:27.542Z" }, + { url = "https://files.pythonhosted.org/packages/1c/7c/fe2980fc03723272e30f135b62360b075f513dfe7cc73aef36c7f04012bd/httptools-0.8.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:a3b7387147361c3fd47a0bde763c5c91b5b4cd4dc9989b8ece84ff436c99843b", size = 463189, upload-time = "2026-05-25T22:17:28.546Z" }, + { url = "https://files.pythonhosted.org/packages/15/1b/47fc5fff68acd1bfa20b4734059c9a06cadb88119dcd5258b5b0d21d91c8/httptools-0.8.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:f256d6ce930c52ca1cb2a960b7da03548c454e7d28b06059ad41bfe789036ce0", size = 466610, upload-time = "2026-05-25T22:17:29.816Z" }, + { url = "https://files.pythonhosted.org/packages/fd/c4/121648f68ce066d7bd762d6b6d97e620847642d38d54f3d90ff11d947629/httptools-0.8.0-cp314-cp314t-macosx_10_13_universal2.whl", hash = "sha256:de1ed58a974e75d56560acc7e7fed01a454994429456f65209789992e41f2568", size = 215023, upload-time = "2026-05-25T22:17:32.401Z" }, + { url = "https://files.pythonhosted.org/packages/b9/b0/312a062ae741ae3e8baa8c8bf20be81b2e67337b259ab4349bebc7b6142e/httptools-0.8.0-cp314-cp314t-macosx_11_0_arm64.whl", hash = "sha256:e93c227b595c6926c1acee96891dd9da4be338cfbe82e5cd3bb9d8dd7dc4ac0b", size = 117405, upload-time = "2026-05-25T22:17:33.742Z" }, + { url = "https://files.pythonhosted.org/packages/fc/37/fccd705f795386bb05bf413012fecff2a33e5aa8c2f069096de3e9fd8702/httptools-0.8.0-cp314-cp314t-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:2a021c3a8e65cc125390d72f59b968afca3bdcaff25bd67965e0a055a14946ca", size = 558497, upload-time = "2026-05-25T22:17:34.732Z" }, + { url = "https://files.pythonhosted.org/packages/bd/39/f172e8003576de35f5ba77ff417cf0e34429d35dc014deef15afa337a72c/httptools-0.8.0-cp314-cp314t-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:48774d39cbb70e2b1f71f88852a3087ae1d3a1eb80482bb48c13067ab080c14f", size = 571585, upload-time = "2026-05-25T22:17:35.813Z" }, + { url = "https://files.pythonhosted.org/packages/3e/b9/f5564760af99f3dbbf3f9104dc00e5da27e96cf433c6bdcf77617f70bf3f/httptools-0.8.0-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:88eead8ec8680a9f146c655bc88445a325bd7921cfd8194c7337e9467282427d", size = 543297, upload-time = "2026-05-25T22:17:37.08Z" }, + { url = "https://files.pythonhosted.org/packages/99/67/8d9f2c313618e161b82f3873188e7196126da1d6e29688df40eb3997c77a/httptools-0.8.0-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:2c032fa028f46871ec7e1fc59fc15e8023eab3e6bbe6ece786a1611719a5d081", size = 539535, upload-time = "2026-05-25T22:17:38.032Z" }, ] [[package]] @@ -4153,15 +4159,15 @@ wheels = [ [[package]] name = "uvicorn" -version = "0.40.0" +version = "0.52.4" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, { name = "h11" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/c3/d1/8f3c683c9561a4e6689dd3b1d345c815f10f86acd044ee1fb9a4dcd0b8c5/uvicorn-0.40.0.tar.gz", hash = "sha256:839676675e87e73694518b5574fd0f24c9d97b46bea16df7b8c05ea1a51071ea", size = 81761, upload-time = "2025-12-21T14:16:22.45Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f2/0f/3f86e61397dd33bf2ccf28188c40db6a740658aeebbbf6e7dbc101a1f487/uvicorn-0.52.4.tar.gz", hash = "sha256:73acfee47a0b133c5de13d219492d62d8a31e935f4fe6e41a232451a15379f86", size = 100627, upload-time = "2026-08-19T06:27:41.821Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/3d/d8/2083a1daa7439a66f3a48589a57d576aa117726762618f6bb09fe3798796/uvicorn-0.40.0-py3-none-any.whl", hash = "sha256:c6c8f55bc8bf13eb6fa9ff87ad62308bbbc33d0b67f84293151efe87e0d5f2ee", size = 68502, upload-time = "2025-12-21T14:16:21.041Z" }, + { url = "https://files.pythonhosted.org/packages/f1/79/4a20b54ab0491485ccd8c077db2d39187c7f12b3e15485d38a7be37c81b4/uvicorn-0.52.4-py3-none-any.whl", hash = "sha256:f86e41a149d7d05a9969337e3946a9c171c06a5d42680896daaba624aeac8da1", size = 79871, upload-time = "2026-08-19T06:27:40.36Z" }, ] [package.optional-dependencies] diff --git a/e2e-tests/fixtures/user.ts b/e2e-tests/fixtures/user.ts index 93bd51cb1e..9871389285 100644 --- a/e2e-tests/fixtures/user.ts +++ b/e2e-tests/fixtures/user.ts @@ -44,7 +44,7 @@ export async function createUser( ): Promise { const password = faker.internet.password(); const response: any = await getClient().post("user/", { - name: faker.name.fullName(), + name: faker.person.fullName(), email: faker.internet.email(), password, language: "en", diff --git a/e2e-tests/package.json b/e2e-tests/package.json index 8deefe85d3..d321cd307c 100644 --- a/e2e-tests/package.json +++ b/e2e-tests/package.json @@ -17,7 +17,7 @@ "codegen": "playwright codegen" }, "dependencies": { - "@faker-js/faker": "7.6.0", + "@faker-js/faker": "10.5.0", "@nuxt/test-utils": "^3.21.0", "@playwright/test": "^1.48.0", "axios": "1.18.0", diff --git a/e2e-tests/yarn.lock b/e2e-tests/yarn.lock index 95d2b24856..0dd4bca9ab 100644 --- a/e2e-tests/yarn.lock +++ b/e2e-tests/yarn.lock @@ -44,10 +44,10 @@ picocolors "^1.0.0" sisteransi "^1.0.5" -"@faker-js/faker@7.6.0": - version "7.6.0" - resolved "https://registry.yarnpkg.com/@faker-js/faker/-/faker-7.6.0.tgz#9ea331766084288634a9247fcd8b84f16ff4ba07" - integrity sha512-XK6BTq1NDMo9Xqw/YkYyGjSsg44fbNwYRx7QK2CuoQgyy+f1rrTDHoExVM5PsyXCtfl2vs2vVJ0MN0yN6LppRw== +"@faker-js/faker@10.5.0": + version "10.5.0" + resolved "https://registry.yarnpkg.com/@faker-js/faker/-/faker-10.5.0.tgz#d2f6a8c7f08d087ac5f077d6babd0821edf24c03" + integrity sha512-bsxD8WLS5lIj7aaoCx1YJkktqYj5vlBUE6HWzu2Q51ksrGJ0H737ECCKlFU7Yf8Br45z9t99frBp/J7kzbMPAg== "@jridgewell/gen-mapping@^0.3.5": version "0.3.13" diff --git a/enterprise/backend/src/baserow_enterprise/assistant/assistant.py b/enterprise/backend/src/baserow_enterprise/assistant/assistant.py index 3c3deb187c..a831f4ef96 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/assistant.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/assistant.py @@ -363,7 +363,10 @@ async def _save_ai_response( await AssistantChatPrediction.objects.acreate( human_message=human_msg, ai_response=ai_msg, - prediction={"answer": answer}, + prediction={ + "answer": answer, + "posthog_trace_id": self._telemetry.trace_id, + }, ) return AiMessage( id=ai_msg.id, @@ -469,7 +472,11 @@ async def _run_agent( """ try: - with self._telemetry.trace(self._chat, user_prompt) as tracer: + with self._telemetry.trace( + self._chat, + user_prompt, + cancelled_by_user=lambda: self._tool_helpers.is_cancelled, + ) as tracer: answer, run_result = await self._run_agent_with_retries( user_prompt, message_history, queue ) @@ -635,10 +642,16 @@ def _looks_like_json_tool_call(text: str) -> bool: """Return True if *text* looks like a tool call dumped as JSON. Checks for ``{"name": ..., "arguments": ...}`` pattern in the first - 200 chars. Does not require valid JSON (the output may be truncated). + 200 chars after an optional code fence. Does not require valid JSON + because the output may be truncated. + + :param text: The final text returned by the agent. + :return: Whether the text appears to contain an unexecuted tool call. """ stripped = text.strip() + if stripped.startswith("```"): + stripped = stripped.split("\n", 1)[-1].strip() return ( bool(stripped) and stripped[0] == "{" diff --git a/enterprise/backend/src/baserow_enterprise/assistant/prompts.py b/enterprise/backend/src/baserow_enterprise/assistant/prompts.py index 31b81e057b..46a7f30cf3 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/prompts.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/prompts.py @@ -8,17 +8,24 @@ """ RULES = """\ + +Three invariants. A tool call that breaks one is invalid — do not send it. +A. IDs. Every `*_id` argument must carry a real ID you have in hand: returned by a tool call in this conversation, present in ``, or given to you by the user. Never invent, guess, or carry over an ID from a different resource. Baserow IDs start at 1, so 0 is never an ID. If you do not have the ID yet, call the list_*/create_* tool that returns it, then pass back the exact value it returned. +B. Modes. Have tools → call them. Each tool is owned by exactly one mode, and `` is the authority: it names what the current `` can call and what each other mode owns. To use a tool owned by another mode, call switch_mode first. If a tool call comes back rejected as an unknown name, that means wrong mode, not missing feature — re-read ``, switch to the owning mode, and retry it once. Only describe manual UI steps once you have confirmed no mode owns a tool for the action; `` lists what genuinely cannot be done in any mode. +C. Payloads. Send every required argument on the first attempt, not only the ones you are confident about. For create_*, update_* and setup_* tools the payload is the point of the call: one carrying just IDs and a `thought` is always incomplete. + 1. Use the `thought` parameter on EVERY tool call. It is shown to the user, so write it as a brief user-facing status (e.g. "Checking existing pages" not "Calling list_pages to get page IDs"). Never use tool names or internal references. -2. Have tools → call them. No tools in current mode → check other modes before saying something is not possible. If another mode has the tool, switch_mode and use it. Only explain manual UI steps if no mode covers the action. -3. One tool per turn. Wait for the result. Never reply and call a tool in same turn. -4. Request priority: action > follow-up (reuse prior IDs, never search docs) > question. When a tool result contains next_steps, act on them immediately — do not ask for permission to continue. -5. You start in the mode matching your UI context (database/application/automation). If the user asks a how-to or feature question, call switch_mode("explain"), then search_user_docs. -6. After finishing the tool calls in a different mode (not just after switching — after the actual work is done and results received), switch back to the original domain mode (check and ). -7. Reply in concise Markdown. Never expose raw JSON or internal IDs unless asked. -8. Before starting work, use list_* to understand what exists and avoid duplicates. But don't list resources you just created — create_* tools already return IDs and refs. When a request references resources by name/ID, verify they exist before building on them. If not found, ask — don't guess. But when the task *requires* creating resources in another domain (e.g. building an app that needs new tables), switch_mode and create them yourself — don't ask the user to do it manually. -9. Before responding to the user, verify ALL parts of `` are addressed. If anything is missing, continue working. -10. At the start, verify the request fits the current UI context (e.g. don't add "Inquiries" table to a "Project Management" DB). If it doesn't match and not explicitly requested, ask the user which target to use. +2. One tool per turn. Wait for the result. Never reply and call a tool in same turn. +3. Request priority: action > follow-up (reuse prior IDs, never search docs) > question. When a tool result contains next_steps, act on them immediately — do not ask for permission to continue. +4. You start in the mode matching your UI context (database/application/automation). If the user asks a how-to or feature question, call switch_mode("explain"), then search_user_docs. +5. After finishing the tool calls in a different mode (not just after switching — after the actual work is done and results received), switch back to the original domain mode (check and ). +6. Reply in concise Markdown. Never expose raw JSON or internal IDs unless asked. +7. Before starting work, use list_* to understand what exists and avoid duplicates. But don't list resources you just created — create_* tools already return IDs and refs. When a request references resources by name/ID, verify they exist before building on them. If not found, ask — don't guess. But when the task *requires* creating resources in another domain (e.g. building an app that needs new tables), switch_mode and create them yourself — don't ask the user to do it manually. +8. Before responding to the user, verify ALL parts of `` are addressed. If anything is missing, continue working. +9. At the start, verify the request fits the current UI context (e.g. don't add "Inquiries" table to a "Project Management" DB). If it doesn't match and not explicitly requested, ask the user which target to use. +10. When a task needs a database, application, or automation that does not exist yet, call create_builders first and build on the ID it returns (contract A). +11. For database formula creation or repair, call generate_formula so the result is validated. Never return or save a handwritten formula. Use save_to_field=true when the user asks to create, fix, save, or apply it; use false only when they explicitly want formula text without changing the table. """ diff --git a/enterprise/backend/src/baserow_enterprise/assistant/retrying_model.py b/enterprise/backend/src/baserow_enterprise/assistant/retrying_model.py index c5da80930a..02c4a75fea 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/retrying_model.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/retrying_model.py @@ -67,6 +67,8 @@ def _is_transient_provider_error(exc: Exception) -> bool: """Return True for provider errors that are transient and safe to retry.""" + if isinstance(exc, ModelHTTPError) and exc.status_code == 429: + return True msg = str(exc) return any(needle in msg for needle in _RETRYABLE_MESSAGES) @@ -413,6 +415,8 @@ async def request( ): raise delay = self._delay_for(attempt) + if isinstance(exc, ModelHTTPError) and exc.retry_after is not None: + delay = min(exc.retry_after, self.max_delay) logger.warning( "[assistant] Model request failed (attempt {}/{}), " "retrying in {:.1f}s: {}", diff --git a/enterprise/backend/src/baserow_enterprise/assistant/telemetry.py b/enterprise/backend/src/baserow_enterprise/assistant/telemetry.py index 820dbc7ee0..6fbf0d098c 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/telemetry.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/telemetry.py @@ -37,6 +37,8 @@ from contextvars import ContextVar from dataclasses import dataclass from datetime import datetime, timezone +from enum import StrEnum +from typing import Callable from uuid import uuid4 from loguru import logger @@ -546,6 +548,21 @@ def _add_phoenix_processors( # --------------------------------------------------------------------------- +class AssistantTraceOutcome(StrEnum): + """How one assistant run ended, recorded on every ``$ai_trace`` event.""" + + ANSWERED = "answered" + ERROR = "error" + CANCELLED = "cancelled" + INTERRUPTED = "interrupted" + NO_ANSWER = "no_answer" + + +_FAILED_OUTCOMES = frozenset( + {AssistantTraceOutcome.ERROR, AssistantTraceOutcome.NO_ANSWER} +) + + class PosthogTracingCallback: """Per-request trace lifecycle. Creates the ``$ai_trace`` event and publishes ``_TraceContext`` for the span exporter.""" @@ -561,11 +578,21 @@ def __init__(self): self.trace_outputs = None @contextmanager - def trace(self, chat: AssistantChat, human_message: str): + def trace( + self, + chat: AssistantChat, + human_message: str, + cancelled_by_user: Callable[[], bool] | None = None, + ): """Context manager that scopes a single assistant execution. Publishes ``_trace_ctx`` so ``PosthogSpanExporter`` can attach trace metadata to child ``$ai_generation`` / ``$ai_span`` events. + + :param chat: The chat whose message this execution answers. + :param human_message: The user message that started this execution. + :param cancelled_by_user: Predicate telling whether the user asked to + stop, used to tell a deliberate cancel apart from a dropped run. """ self.chat = chat @@ -589,23 +616,25 @@ def trace(self, chat: AssistantChat, human_message: str): ) tools_token = _tool_calls.set([]) - exception = None + exception: Exception | None = None + interruption: BaseException | None = None try: yield self except Exception as exc: exception = exc raise + except BaseException as exc: + interruption = exc + raise finally: tool_call_names = _tool_calls.get([]) _trace_ctx.reset(token) _tool_calls.reset(tools_token) - output_state = self.trace_outputs if exception is None else str(exception) - if tool_call_names: - if output_state is None: - output_state = {} - if isinstance(output_state, dict): - output_state["tool_calls"] = tool_call_names + outcome = self._resolve_outcome(exception, interruption, cancelled_by_user) + output_state = self._build_output_state(outcome, exception) + if tool_call_names and isinstance(output_state, dict): + output_state["tool_calls"] = tool_call_names self._capture_event( "$ai_trace", @@ -615,9 +644,10 @@ def trace(self, chat: AssistantChat, human_message: str): "$ai_span_name": f"{self.user_id}: {human_message[:20]}", "$ai_span_id": self.span_id, "$ai_latency": (_utc_now() - start_time).total_seconds(), - "$ai_is_error": exception is not None, + "$ai_is_error": outcome in _FAILED_OUTCOMES, "$ai_input_state": {"user_message": human_message}, "$ai_output_state": output_state, + "assistant_outcome": outcome.value, }, ) @@ -626,6 +656,48 @@ def trace(self, chat: AssistantChat, human_message: str): except Exception: pass + def _resolve_outcome( + self, + exception: Exception | None, + interruption: BaseException | None, + cancelled_by_user: Callable[[], bool] | None, + ) -> AssistantTraceOutcome: + """Classify how a single assistant execution ended. + + :param exception: The error raised inside the traced block, if any. + :param interruption: The ``BaseException`` that unwound the traced + block, if any. + :param cancelled_by_user: Predicate telling whether the user asked to + stop. + :return: The outcome to record on the ``$ai_trace`` event. + """ + + if exception is not None: + return AssistantTraceOutcome.ERROR + if interruption is not None: + if cancelled_by_user is not None and cancelled_by_user(): + return AssistantTraceOutcome.CANCELLED + return AssistantTraceOutcome.INTERRUPTED + if self.trace_outputs is None: + return AssistantTraceOutcome.NO_ANSWER + return AssistantTraceOutcome.ANSWERED + + def _build_output_state( + self, outcome: AssistantTraceOutcome, exception: Exception | None + ) -> dict | str: + """Render the ``$ai_output_state`` payload for *outcome*. + + :param outcome: The classified outcome of the execution. + :param exception: The error raised inside the traced block, if any. + :return: The answer, the error message, or the terminal status. + """ + + if outcome is AssistantTraceOutcome.ERROR: + return str(exception) + if outcome is AssistantTraceOutcome.ANSWERED: + return self.trace_outputs + return {"status": outcome.value} + def set_trace_output(self, output: str): """Record the agent's final answer for the ``$ai_trace`` event.""" diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/agents.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/agents.py index 9bbaff7546..d39467a588 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/agents.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/agents.py @@ -137,16 +137,16 @@ def _generate_node_formulas(node: ActionNodeCreate, orm_node: AutomationNode): def update_single_node_formulas( - node_update: "NodeUpdate", + node: "NodeUpdate | ActionNodeCreate", orm_node: AutomationNode, tool_helpers: "ToolHelpers", ) -> None: - """Generate and apply formulas for one node being updated. + """Generate and apply formulas for one node being created or updated. Builds formula context from the node's workflow, then generates - formulas for the $formula: fields in the update. + formulas for the $formula: fields in the payload. - :param node_update: The assistant node update containing formula requests. + :param node: The assistant node creation or update containing formula requests. :param orm_node: The persisted automation node to update. :param tool_helpers: Helpers for status updates and the request model profile. :return: None. @@ -165,10 +165,14 @@ def update_single_node_formulas( metadata["node_id"] = wf_node.id context.add_node_context(wf_node.id, example, metadata) - formulas_to_create = node_update.get_formulas_to_update(orm_node) + formulas_to_create = ( + node.get_formulas_to_update(orm_node) + if isinstance(node, NodeUpdate) + else node.get_formulas_to_create(orm_node) + ) if formulas_to_create is None: return result = generate_formula(formulas_to_create, context) if result: - node_update.update_service_with_formulas(orm_node.service, result) + node.update_service_with_formulas(orm_node.service, result) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/helpers.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/helpers.py index 82c7493621..4c58b071d4 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/helpers.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/helpers.py @@ -265,13 +265,20 @@ def update_node( node_type = node.service.get_type().type if node.service else None service_dict = node_update.to_update_service_dict(node_type) if node_type else None if service_dict is not None: + if ( + node_type == "local_baserow_upsert_row" + and "table_id" in service_dict + and service_dict["table_id"] != node.service.specific.table_id + ): + # Mappings from the previous table cannot be dispatched on the new one. + service_dict["field_mappings"] = [] kwargs["service"] = service_dict - if kwargs: - tool_helpers.update_status( - _("Updating node '%(label)s'..." % {"label": node.label}) - ) - AutomationNodeService().update_node(user, node.id, **kwargs) + tool_helpers.update_status( + _("Updating node '%(label)s'..." % {"label": node.label}) + ) + # Deferred row values still need update permission and test-clone invalidation. + AutomationNodeService().update_node(user, node.id, **kwargs) return AutomationNodeService().get_node(user, node_update.node_id) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/tools.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/tools.py index 9db61c34aa..799cfff2d9 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/tools.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/tools.py @@ -3,12 +3,14 @@ from django.db import transaction from django.utils.translation import gettext as _ +from loguru import logger from pydantic import Field from pydantic_ai import RunContext from pydantic_ai.toolsets import FunctionToolset from baserow.contrib.automation.workflows.service import AutomationWorkflowService from baserow_enterprise.assistant.deps import AssistantDeps +from baserow_enterprise.assistant.tools.shared import require_payload from baserow_enterprise.assistant.types import WorkflowNavigationType from . import agents, helpers @@ -100,14 +102,14 @@ def add_nodes( RETURNS: Created nodes array with id, label, type. DO NOT USE when: You want to create an entirely new workflow — use create_workflows instead. HOW: Use list_nodes first to find the existing node IDs, then specify previous_node_ref to place new nodes. Use router_edge_label when attaching to a router branch. + REQUIRED: `workflow_id` and `nodes` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. """ user = ctx.deps.user workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not nodes: - return {"created_nodes": []} + require_payload("add_nodes", "nodes", nodes) tool_helpers.update_status(_("Adding nodes to workflow...")) @@ -120,9 +122,9 @@ def add_nodes( # Generate formulas for nodes that need them for orm_node, node_create in [(n, nodes[i]) for i, n in enumerate(created_nodes)]: + node_create.apply_direct_values(orm_node.service) formulas = node_create.get_formulas_to_create(orm_node) if formulas: - node_create.apply_direct_values(orm_node.service) tool_helpers.update_status( _( "Generating formulas for node '%(label)s'..." @@ -135,8 +137,6 @@ def add_nodes( node_create, orm_node, tool_helpers ) except Exception: - from loguru import logger - logger.exception( "Failed to generate formulas for node {}", orm_node.id ) @@ -193,8 +193,7 @@ def create_workflows( workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not workflows: - return {"created_workflows": []} + require_payload("create_workflows", "workflows", workflows) created = [] @@ -267,7 +266,7 @@ def update_nodes( updated = [] errors = [] - nodes_needing_formulas = [] + updated_pairs = [] with transaction.atomic(): for node_update in nodes: @@ -277,17 +276,16 @@ def update_nodes( user, workspace, node_update, tool_helpers ) updated.append({"node_id": orm_node.id, "label": orm_node.label}) - - # Check if any fields need formula generation - formulas = node_update.get_formulas_to_update(orm_node) - if formulas: - nodes_needing_formulas.append((node_update, orm_node, formulas)) + updated_pairs.append((node_update, orm_node)) except Exception as e: errors.append(f"Error updating node {node_update.node_id}: {e}") # Apply direct values and generate formulas outside the main transaction - for node_update, orm_node, formulas in nodes_needing_formulas: + for node_update, orm_node in updated_pairs: + # Literal values must be applied whether or not any formula follows. node_update.apply_direct_values(orm_node.service) + if not node_update.get_formulas_to_update(orm_node): + continue tool_helpers.update_status( _("Generating formulas for node '%(label)s'..." % {"label": orm_node.label}) ) @@ -295,8 +293,6 @@ def update_nodes( try: agents.update_single_node_formulas(node_update, orm_node, tool_helpers) except Exception as exc: - from loguru import logger - logger.exception( "Failed to generate formulas for node {}: {}", orm_node.id, exc ) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/types/node.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/types/node.py index 8c3e98ccf2..a0cdaf0c86 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/types/node.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/automation/types/node.py @@ -140,17 +140,69 @@ class AutomationFieldValue(BaseModel): _PERIODIC_KEYS = {"interval", "minute", "hour", "day_of_week", "day_of_month"} +# list_nodes reports registered names, but dispatch tables key on short aliases. +CANONICAL_TO_SHORT_TYPE = { + "local_baserow_rows_created": "rows_created", + "local_baserow_rows_updated": "rows_updated", + "local_baserow_rows_deleted": "rows_deleted", + "local_baserow_create_row": "create_row", + "local_baserow_update_row": "update_row", + "local_baserow_delete_row": "delete_row", +} + + +def _fold_type_alias(data): + """Normalize node aliases while leaving malformed types for validation. + + :param data: The raw node payload before model validation. + :return: The payload with any registered type replaced by its short alias. + """ + + if isinstance(data, dict): + node_type = data.get("type") + if isinstance(node_type, str) and node_type in CANONICAL_TO_SHORT_TYPE: + data["type"] = CANONICAL_TO_SHORT_TYPE[node_type] + return data + + +# Row-action service types differ from the short node names the tables key on. +_SERVICE_TO_DISPATCH_TYPE = { + "local_baserow_upsert_row": "update_row", + "local_baserow_delete_row": "delete_row", +} + + +def _service_dispatch_type(service: "Service | None") -> str | None: + if service is None: + return None + service_type = service.get_type().type + return _SERVICE_TO_DISPATCH_TYPE.get(service_type, service_type) + + +ROW_TRIGGER_TYPES = frozenset( + { + "rows_created", + "rows_updated", + "rows_deleted", + } +) + + class TriggerNodeCreate(BaseModel): """Create a trigger node in a workflow.""" ref: str = Field(..., description="Temporary reference ID for creation.") label: str = Field(..., description="Display name.") + # Registered names are accepted too: the model echoes what list_nodes returned. type: Literal[ "periodic", "http_trigger", "rows_updated", "rows_created", "rows_deleted", + "local_baserow_rows_updated", + "local_baserow_rows_created", + "local_baserow_rows_deleted", ] periodic_interval: Optional[PeriodicTriggerSettings] = Field( @@ -162,6 +214,11 @@ class TriggerNodeCreate(BaseModel): description="(rows_*) Table to monitor.", ) + @model_validator(mode="before") + @classmethod + def _fold_registered_type(cls, data): + return _fold_type_alias(data) + @model_validator(mode="before") @classmethod def _fold_flat_periodic(cls, data): @@ -180,7 +237,7 @@ def _fold_flat_periodic(cls, data): def _validate_trigger_settings(self): if self.type == "periodic" and self.periodic_interval is None: raise ValueError("periodic trigger requires periodic_interval") - if self.type in ("rows_created", "rows_updated", "rows_deleted"): + if self.type in ROW_TRIGGER_TYPES: if self.rows_triggers_settings is None: raise ValueError(f"{self.type} trigger requires rows_triggers_settings") return self @@ -197,10 +254,7 @@ def to_orm_service_dict(self) -> dict[str, Any]: ) return values - if ( - self.type in ["rows_created", "rows_updated", "rows_deleted"] - and self.rows_triggers_settings - ): + if self.type in ROW_TRIGGER_TYPES and self.rows_triggers_settings: return self.rows_triggers_settings.model_dump() return {} @@ -219,6 +273,7 @@ class TriggerNodeItem(TriggerNodeCreate): # Action node # --------------------------------------------------------------------------- +# Registered names are accepted too: the model echoes what list_nodes returned. ActionNodeType = Literal[ "router", "smtp_email", @@ -227,12 +282,20 @@ class TriggerNodeItem(TriggerNodeCreate): "update_row", "delete_row", "ai_agent", + "local_baserow_create_row", + "local_baserow_update_row", + "local_baserow_delete_row", ] class ActionNodeCreate(BaseModel): """Flat model for creating an action node: type + type-specific fields.""" + @model_validator(mode="before") + @classmethod + def _fold_registered_type(cls, data): + return _fold_type_alias(data) + ref: str = Field(..., description="Temporary reference ID for creation.") label: str = Field(..., description="Display name.") type: ActionNodeType @@ -372,11 +435,12 @@ def apply_direct_values(self, service: Service): fn = _APPLY_DIRECT.get(self.type) if fn is not None: - fn(self, service) + fn(self, service.specific) def update_service_with_formulas(self, service: Service, formulas: dict[str, str]): """Write generated formulas back to the ORM service.""" + service = service.specific fn = _UPDATE_FORMULAS.get(self.type) if fn is not None: fn(self, service, formulas) @@ -395,9 +459,9 @@ def _router_to_orm(n: ActionNodeCreate) -> dict[str, Any]: def _email_to_orm(n: ActionNodeCreate) -> dict[str, Any]: return { - "to_email": literal_or_placeholder(n.to_emails), - "cc_email": literal_or_placeholder(n.cc_emails), - "bcc_email": literal_or_placeholder(n.bcc_emails), + "to_emails": literal_or_placeholder(n.to_emails), + "cc_emails": literal_or_placeholder(n.cc_emails), + "bcc_emails": literal_or_placeholder(n.bcc_emails), "subject": literal_or_placeholder(n.subject), "body": literal_or_placeholder(n.body), "body_type": f"'{n.body_type}'", @@ -673,7 +737,9 @@ class NodeUpdate(BaseModel): def to_update_service_dict(self, current_type: str) -> dict[str, Any] | None: """Build a service kwargs dict from non-None fields. Returns None if no service fields set.""" - builder = _TO_UPDATE_SERVICE.get(current_type) + builder = _TO_UPDATE_SERVICE.get( + _SERVICE_TO_DISPATCH_TYPE.get(current_type, current_type) + ) if builder is None: return None result = builder(self) @@ -681,20 +747,19 @@ def to_update_service_dict(self, current_type: str) -> dict[str, Any] | None: def get_formulas_to_update(self, orm_node: AutomationNode) -> dict[str, str] | None: """Return a {key: description} dict of formulas to generate, or None.""" - fn = _GET_UPDATE_FORMULAS.get( - orm_node.service.get_type().type if orm_node.service else None - ) + fn = _GET_UPDATE_FORMULAS.get(_service_dispatch_type(orm_node.service)) return fn(self, orm_node) if fn else None def apply_direct_values(self, service: Service): """Apply literal (non-$formula) values directly to the service.""" - fn = _APPLY_UPDATE_DIRECT.get(service.get_type().type if service else None) + fn = _APPLY_UPDATE_DIRECT.get(_service_dispatch_type(service)) if fn is not None: - fn(self, service) + fn(self, service.specific) def update_service_with_formulas(self, service: Service, formulas: dict[str, str]): """Write generated formulas back to the ORM service.""" - stype = service.get_type().type if service else None + stype = _service_dispatch_type(service) + service = service.specific fn = _UPDATE_FORMULAS.get(stype) if fn is not None: # Reuse the existing dispatch (expects ActionNodeCreate-like but works for our purposes) @@ -709,11 +774,11 @@ def update_service_with_formulas(self, service: Service, formulas: dict[str, str def _email_update_service(n: "NodeUpdate") -> dict[str, Any]: d = {} if n.to_emails is not None: - d["to_email"] = literal_or_placeholder(n.to_emails) + d["to_emails"] = literal_or_placeholder(n.to_emails) if n.cc_emails is not None: - d["cc_email"] = literal_or_placeholder(n.cc_emails) + d["cc_emails"] = literal_or_placeholder(n.cc_emails) if n.bcc_emails is not None: - d["bcc_email"] = literal_or_placeholder(n.bcc_emails) + d["bcc_emails"] = literal_or_placeholder(n.bcc_emails) if n.subject is not None: d["subject"] = literal_or_placeholder(n.subject) if n.body is not None: @@ -791,7 +856,7 @@ def _slack_update_formulas( return None -def _row_action_update_formulas( +def _get_row_action_update_formulas( n: "NodeUpdate", orm_node: AutomationNode ) -> dict[str, str] | None: from baserow_enterprise.assistant.tools.shared.formula_utils import ( @@ -828,9 +893,9 @@ def _ai_agent_update_formulas( _GET_UPDATE_FORMULAS: dict[str, Callable] = { "smtp_email": _email_update_formulas, "slack_write_message": _slack_update_formulas, - "create_row": _row_action_update_formulas, - "update_row": _row_action_update_formulas, - "delete_row": _row_action_update_formulas, + "create_row": _get_row_action_update_formulas, + "update_row": _get_row_action_update_formulas, + "delete_row": _get_row_action_update_formulas, "ai_agent": _ai_agent_update_formulas, } diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/tools.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/tools.py index 394a2b295d..f1d9d5f94f 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/tools.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/tools.py @@ -21,6 +21,7 @@ ThemeName, apply_theme, ) +from baserow_enterprise.assistant.tools.shared import require_payload from baserow_enterprise.assistant.types import BuilderPageNavigationType from . import agents, helpers @@ -140,6 +141,7 @@ def create_pages( WHAT it does: Creates pages with paths and parameters. Skips duplicates by name. RETURNS: Created pages with id, name, path. DO NOT USE when: Pages with those names already exist — check with list_pages first. + REQUIRED: `application_id` and `pages` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Page Setup - Each page needs a unique name and path. @@ -154,8 +156,7 @@ def create_pages( workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not pages: - return {"created_pages": []} + require_payload("create_pages", "pages", pages) builder = helpers.get_builder(user, workspace, application_id) @@ -293,6 +294,7 @@ def create_data_sources( WHAT it does: Creates list_rows or get_row data sources. Skips duplicates by name. RETURNS: Created data sources with ref-to-ID mapping. DO NOT USE when: Data sources with those names already exist on the page. + REQUIRED: `page_id` and `data_sources` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Data Source Types - list_rows: Fetches multiple rows — use when the page displays a collection (table, repeat, dropdown). @@ -311,8 +313,7 @@ def create_data_sources( user = ctx.deps.user tool_helpers = ctx.deps.tool_helpers - if not data_sources: - return {"created_data_sources": [], "ref_to_id_map": {}} + require_payload("create_data_sources", "data_sources", data_sources) page = helpers.get_page(user, page_id) integration = helpers.get_local_baserow_integration(user, page.builder) @@ -444,7 +445,7 @@ def list_elements( List all elements on a page. WHEN to use: Check existing elements, find element IDs or container structure. - WHAT it does: Lists elements with id, type, order, parent_element_id, is_container. + WHAT it does: Lists elements with id, type, parent_element_id, is_container. RETURNS: Elements array. Elements with page_name="[shared]" are headers/footers visible on ALL pages. @@ -468,14 +469,15 @@ def _create_elements_internal( page_id: int, elements: list[ElementItemCreate], before_element_id: int | None = None, + *, + tool_name: str, ) -> dict[str, Any]: """Shared implementation for all create_*_elements tools.""" user = ctx.deps.user tool_helpers = ctx.deps.tool_helpers - if not elements: - return {"created_elements": [], "ref_to_id_map": {}} + require_payload(tool_name, "elements", elements) page = helpers.get_page(user, page_id) shared_page = PageHandler().get_shared_page(page.builder) @@ -588,6 +590,7 @@ def create_display_elements( WHEN to use: User wants to add text content, headings, buttons, links, or images. WHAT it does: Creates display elements with formula support for dynamic values. RETURNS: Created elements with ref-to-ID mapping. + REQUIRED: `page_id` and `elements` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Element Structure - parent_element: int ID (existing container) or string ref (same batch) @@ -609,7 +612,13 @@ def create_display_elements( """ internal = [el.to_element_item_create() for el in elements] - return _create_elements_internal(ctx, page_id, internal, before_element_id) + return _create_elements_internal( + ctx, + page_id, + internal, + before_element_id, + tool_name="create_display_elements", + ) def create_layout_elements( @@ -632,6 +641,7 @@ def create_layout_elements( WHEN to use: User wants page structure — columns, containers, headers, footers, menus. WHAT it does: Creates container elements that hold child elements. RETURNS: Created elements with ref-to-ID mapping. + REQUIRED: `page_id` and `elements` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Element Structure - Layout elements are containers — add children via parent_element ref. @@ -649,7 +659,13 @@ def create_layout_elements( """ internal = [el.to_element_item_create() for el in elements] - return _create_elements_internal(ctx, page_id, internal, before_element_id) + return _create_elements_internal( + ctx, + page_id, + internal, + before_element_id, + tool_name="create_layout_elements", + ) def create_form_elements( @@ -672,6 +688,7 @@ def create_form_elements( WHEN to use: User wants a form to collect input data. WHAT it does: Creates form containers and input elements with validation. RETURNS: Created elements with ref-to-ID mapping. + REQUIRED: `page_id` and `elements` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Form Structure - Create a form_container first, then add inputs inside it via parent_element. @@ -690,7 +707,13 @@ def create_form_elements( """ internal = [el.to_element_item_create() for el in elements] - return _create_elements_internal(ctx, page_id, internal, before_element_id) + return _create_elements_internal( + ctx, + page_id, + internal, + before_element_id, + tool_name="create_form_elements", + ) def create_collection_elements( @@ -714,6 +737,7 @@ def create_collection_elements( WHEN to use: User wants to display data from a data source in a table or repeating layout. WHAT it does: Creates collection elements connected to data sources. RETURNS: Created elements with ref-to-ID mapping. + REQUIRED: `page_id` and `elements` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Prerequisites Create data sources first (create_data_sources), then reference them here. @@ -729,7 +753,13 @@ def create_collection_elements( """ internal = [el.to_element_item_create() for el in elements] - return _create_elements_internal(ctx, page_id, internal, before_element_id) + return _create_elements_internal( + ctx, + page_id, + internal, + before_element_id, + tool_name="create_collection_elements", + ) # --------------------------------------------------------------------------- @@ -926,7 +956,6 @@ def move_elements( "element_id": element.id, "parent_element_id": element.parent_element_id, "place_in_container": element.place_in_container, - "order": str(element.order), } ) except Exception as exc: @@ -986,6 +1015,7 @@ def create_actions( WHEN to use: User wants buttons/forms to perform actions (navigate, create/update rows, show notifications). WHAT it does: Creates workflow actions with formula support for dynamic values. RETURNS: Created actions with id, type, element_ref, event. + REQUIRED: `page_id` and `actions` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. ## Attaching Actions - element_ref: attach to newly created element (auto-tracked) @@ -1010,8 +1040,7 @@ def create_actions( user = ctx.deps.user tool_helpers = ctx.deps.tool_helpers - if not actions: - return {"created_actions": []} + require_payload("create_actions", "actions", actions) page = helpers.get_page(user, page_id) integration = helpers.get_local_baserow_integration(user, page.builder) @@ -1298,6 +1327,7 @@ def setup_page( WHEN to use: Building a complete page with data, UI elements, and interactions. WHAT it does: Creates data sources first, then elements (in order), then actions. Handles ref resolution across all three phases. RETURNS: Created items with ref-to-ID mappings and any errors. Partial success is possible — some items may be created even when others fail. Check the ``errors`` key. + ARGS: ``page_id`` is WHERE to build; ``data_sources``, ``elements`` and ``actions`` are WHAT to build. Each of the three is optional on its own, but a call carrying none of them creates nothing and is rejected. To only open a page, use ``navigate``. ## Deduplication Data sources are deduplicated by name (case-insensitive) and by structural match (same type and table). Existing data sources are reused and their IDs mapped to the provided refs. @@ -1343,6 +1373,15 @@ def setup_page( user = ctx.deps.user tool_helpers = ctx.deps.tool_helpers + if not data_sources and not elements and not actions: + raise helpers.ToolInputError( + "setup_page was called with no content: `data_sources`, `elements` " + "and `actions` were all empty, so nothing was created. An ID " + "argument says where to act, never what to do — send the items to " + "build in the same call that carries the ID. To only open a page, " + "use `navigate` instead." + ) + page = helpers.get_page(user, page_id) shared_page = PageHandler().get_shared_page(page.builder) integration = helpers.get_local_baserow_integration(user, page.builder) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/data_source.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/data_source.py index ec0c13db8d..52126dc204 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/data_source.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/data_source.py @@ -27,7 +27,18 @@ # Data source sort # --------------------------------------------------------------------------- -DataSourceType = Literal["list_rows", "get_row"] +DataSourceType = Literal[ + "list_rows", + "get_row", + "local_baserow_list_rows", + "local_baserow_get_row", +] + +# list_data_sources reports registered names; the tables below key on short forms. +_CANONICAL_TO_SHORT_TYPE = { + "local_baserow_list_rows": "list_rows", + "local_baserow_get_row": "get_row", +} class DataSourceSort(BaseModel): @@ -76,6 +87,21 @@ class DataSourceCreate(BaseModel): the correct required fields per type. """ + @model_validator(mode="before") + @classmethod + def _fold_registered_type(cls, data): + """Normalize source aliases while leaving malformed types for validation. + + :param data: The raw data source payload before model validation. + :return: The payload with any registered type replaced by its short alias. + """ + + if isinstance(data, dict): + source_type = data.get("type") + if isinstance(source_type, str) and source_type in _CANONICAL_TO_SHORT_TYPE: + data["type"] = _CANONICAL_TO_SHORT_TYPE[source_type] + return data + ref: str = Field(..., description="Reference ID for this data source.") name: str = Field(..., description="Human-readable name.") type: DataSourceType = Field(..., description="'list_rows' or 'get_row'.") @@ -126,7 +152,8 @@ def matches_existing(self, existing: "DataSourceItem") -> bool: Delegates to a per-type matcher in ``_STRUCTURAL_MATCH``. """ - if self.type != existing.type: + existing_type = _CANONICAL_TO_SHORT_TYPE.get(existing.type, existing.type) + if self.type != existing_type: return False matcher = _STRUCTURAL_MATCH.get(self.type) return matcher(self, existing) if matcher else False diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/element.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/element.py index a9a5d85dd5..3550b14ce1 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/element.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/element.py @@ -117,14 +117,21 @@ class MenuItemCreate(BaseModel): class TableFieldConfig(BaseModel): """ - Column configuration for table elements. + One column of a table element. - ``type`` is ``"text"`` (default) or ``"button"``. + - type="text" (default): ``value`` is the cell content — a literal or + "$formula: ". Omit ``value`` to show the data source field + whose name matches ``name``. + - type="button": ``label`` is the button caption. + + These are the only column types, and ``name``, ``type``, ``value``, + ``label`` are the only accepted keys; any other key is rejected. Column + keys are not interchangeable with element keys. """ name: str = Field(..., description="Column header name.") - type: Literal["text", "button", "link", "tags"] = Field( - default="text", description="Column type." + type: Literal["text", "button"] = Field( + default="text", description="Column type: 'text' (default) or 'button'." ) # text columns @@ -806,10 +813,6 @@ def _convert_table_fields( }, } ) - elif field_cfg.type == "link": - result.append({"name": field_cfg.name, "type": "link", "config": {}}) - elif field_cfg.type == "tags": - result.append({"name": field_cfg.name, "type": "tags", "config": {}}) return result @@ -1872,7 +1875,6 @@ class ElementItem(BaseModel): id: int type: str - order: str parent_element_id: int | None = None place_in_container: str | None = None is_container: bool = Field( @@ -1912,7 +1914,6 @@ def from_orm(cls, element) -> "ElementItem": return cls( id=element.id, type=element_type, - order=str(element.order), parent_element_id=element.parent_element_id, place_in_container=element.place_in_container, is_container=element_type in CONTAINER_ELEMENT_TYPES, diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/workflow_action.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/workflow_action.py index 44f3ec07c5..8af1210e0d 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/workflow_action.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/builder/types/workflow_action.py @@ -66,6 +66,14 @@ class FieldValueMapping(BaseModel): "delete_row": ("table_id", "row_id"), } +# Fields that carry a database ID when given as an int; refs stay strings. +_ID_FIELDS: tuple[str, ...] = ( + "element", + "navigate_to_page_id", + "table_id", + "data_source", +) + def _strip_formula_prefix(value: str) -> str: """ @@ -331,10 +339,29 @@ def _update_open_page_formulas( class ActionCreate(BaseModel): """ - Flat model for creating a workflow action. - - All type-specific fields are optional — a ``@model_validator`` - enforces the correct required fields per type. + One workflow action, attached to an element (event "click") or to a form + container (event "submit"). + + ``type`` and ``element`` are always required. Each type also needs: + - notification: title + - open_page: navigate_to_page_id + - create_row: table_id, field_values + - update_row: table_id, row_id, field_values + - delete_row: table_id, row_id + - refresh_data_source: data_source + - logout: nothing more + + Send all of them in the same call: a call missing any of them is rejected + in full, and every retry costs a round trip. Numeric IDs must come from a + tool result; string refs must name an item created in this same call. + + The tag in parentheses on a field description names the types that field + applies to; leave every other field unset. No key outside this model is + accepted. + + open_page navigates to a page of this application only — no action opens + an external URL. External navigation is a property of the link element + (navigation_type, navigate_to_url, link_variant), not of an action. """ type: ActionType = Field(..., description="Action type.") @@ -388,12 +415,30 @@ class ActionCreate(BaseModel): ) @model_validator(mode="after") - def _check_required(self): - for field_name in _REQUIRED_FIELDS.get(self.type, ()): - value = getattr(self, field_name) - # A blank row ID is no row at all: the action would fail every click. - if value is None or (field_name == "row_id" and not value.strip()): - raise ValueError(f"'{field_name}' is required for type '{self.type}'.") + def _check_required(self) -> "ActionCreate": + required = _REQUIRED_FIELDS.get(self.type, ()) + missing = [ + name + for name in required + if getattr(self, name) is None + or (name == "row_id" and not getattr(self, name).strip()) + ] + if missing: + raise ValueError( + f"Type '{self.type}' requires all of: {', '.join(required)}. " + f"Missing: {', '.join(missing)}. Send every required field in a " + f"single call. If a value is not known yet, create or look up the " + f"resource it refers to first, or choose a type that does not " + f"require it." + ) + for name in _ID_FIELDS: + value = getattr(self, name, None) + if isinstance(value, int) and value <= 0: + raise ValueError( + f"'{name}' is {value}, which is never a valid ID. Use an ID " + f"returned by an earlier tool result, or a string ref to an " + f"item created in this same call." + ) return self # -- ORM helpers -------------------------------------------------------- diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/core/tools.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/core/tools.py index 2b696d958b..fe1d97dcef 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/core/tools.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/core/tools.py @@ -172,8 +172,15 @@ def switch_mode( def update_builder( ctx: RunContext[AssistantDeps], + builder_id: Annotated[ + int, + Field( + description="ID of the builder to update, as returned by list_builders or create_builders." + ), + ], update: Annotated[ - BuilderUpdate, Field(description="Application settings to update.") + BuilderUpdate, + Field(description="Settings to change. Set only the fields you want changed."), ], thought: Annotated[ str, Field(description="Brief reasoning for calling this tool.") @@ -192,7 +199,7 @@ def update_builder( user = ctx.deps.user - app = CoreService().get_application(user, update.builder_id).specific + app = CoreService().get_application(user, builder_id).specific ctx.deps.tool_helpers.update_status( _("Updating %(app_name)s...") % {"app_name": app.name} ) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/core/types.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/core/types.py index 80eb74d58f..11bf4c779d 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/core/types.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/core/types.py @@ -190,7 +190,6 @@ class BuilderUpdate(BaseModel): Fields are type-specific — only set the ones relevant to the application type. """ - builder_id: int = Field(..., description="ID of the application to update.") name: str | None = Field(default=None, description="New name.") # Application (builder) specific diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/agents.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/agents.py index 7f43ae8dce..4115700c4b 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/agents.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/agents.py @@ -1,11 +1,21 @@ -from typing import Any, Callable +import re +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any, Callable from django.contrib.auth.models import AbstractUser from django.utils.translation import gettext as _ +from loguru import logger from pydantic import BaseModel as PydanticBaseModel -from pydantic import Field -from pydantic_ai import Agent, Tool +from pydantic import Field, TypeAdapter +from pydantic_ai import Agent, ModelRetry, RunContext, Tool +from pydantic_ai.agent import AgentRunResult +from pydantic_ai.messages import ( + ModelMessage, + RetryPromptPart, + ToolCallPart, + ToolReturnPart, +) from pydantic_ai.toolsets import FunctionToolset from pydantic_ai.usage import UsageLimits @@ -24,6 +34,11 @@ format_sample_rows_prompt, ) +if TYPE_CHECKING: + from baserow_enterprise.assistant.model_profiles import ( + ResolvedAssistantModelProfile, + ) + # --------------------------------------------------------------------------- # Formula generation agent # --------------------------------------------------------------------------- @@ -34,9 +49,11 @@ class FormulaGenerationResult(PydanticBaseModel): table_id: int = Field( description=( - "The ID of the table the formula is intended for. " - "Should be the same as current_table_id, unless the formula can " - "only be created in a different table." + "The ID of the table the formula field belongs to. It must be the " + "`id` of one of the tables in the schema given in the prompt: the " + "table that holds the fields the formula reads directly (for a " + "lookup, the table holding the link field, not the linked table). " + "Never invent an ID, and never use 0 or null." ) ) field_name: str = Field( @@ -60,20 +77,194 @@ class FormulaGenerationResult(PydanticBaseModel): ) +# The default budget is one attempt: the first validator rejection would end it. +FORMULA_AGENT_RETRIES = 3 +FORMULA_AGENT_REQUEST_LIMIT = 20 + formula_generation_agent: Agent[None, FormulaGenerationResult] = Agent( output_type=FormulaGenerationResult, instructions=FORMULA_AGENT_INSTRUCTIONS, name="formula_generation_agent", # Stop as soon as the output tool returns; don't run trailing tool calls. end_strategy="early", + retries=FORMULA_AGENT_RETRIES, ) +GET_FORMULA_TYPE_TOOL_NAME = "get_formula_type" +_TABLE_ID_ADAPTER = TypeAdapter(int) + +# One rejected candidate is not evidence that the language cannot express a request. +FORMULA_MIN_ATTEMPTS_BEFORE_IMPOSSIBLE = 2 + + +def _normalize_formula(formula: str) -> str: + # Used only to discount repetitive attempts, never to prove validity: + # collapsing whitespace can change string literals and quoted field names. + return " ".join(formula.split()) + + +def _formula_attempts( + messages: Sequence[ModelMessage], +) -> tuple[set[tuple[int, str, str]], set[str]]: + """ + Split the formulas passed to get_formula_type during a run into accepted and + rejected. + + A rejection reaches the model as a RetryPromptPart and an acceptance as a + ToolReturnPart, so the two sets are the run's own record of what was checked. + + :param messages: The run's message history. + :returns: Exact accepted (table ID, field name, formula) tuples, and normalized + rejected formulas used to count distinct attempts. + """ + + attempted: dict[str, dict] = {} + accepted_ids: set[str] = set() + rejected_ids: set[str] = set() + for message in messages: + for part in message.parts: + if getattr(part, "tool_name", None) != GET_FORMULA_TYPE_TOOL_NAME: + continue + if isinstance(part, ToolCallPart): + args = part.args_as_dict() + formula = args.get("formula") + if isinstance(formula, str) and formula.strip(): + attempted[part.tool_call_id] = args + elif isinstance(part, RetryPromptPart): + rejected_ids.add(part.tool_call_id) + elif isinstance(part, ToolReturnPart): + accepted_ids.add(part.tool_call_id) + + # History retains raw arguments, including numeric strings the tool coerced. + # Use the same integer validation as the tool when identifying accepted calls. + accepted = { + ( + _TABLE_ID_ADAPTER.validate_python(args["table_id"]), + args["field_name"], + args["formula"], + ) + for call_id, args in attempted.items() + if call_id in accepted_ids + } + rejected = { + _normalize_formula(args["formula"]) + for call_id, args in attempted.items() + if call_id in rejected_ids + } + return accepted, rejected + + +@formula_generation_agent.output_validator +def _verdict_must_be_backed_by_validation( + ctx: RunContext[None], output: FormulaGenerationResult +) -> FormulaGenerationResult: + """ + Send back any verdict get_formula_type did not actually produce. + + Both directions are enforced: a valid verdict must name a formula the tool + accepted for that table and field name, and an impossible verdict must follow + multiple distinct candidates checked by the tool. A verdict with no tool call + behind it is a guess. + + :param ctx: The agent run context. + :param output: The candidate result to validate. + :returns: The output unchanged when it is backed by validation. + :raises ModelRetry: When the verdict is not grounded in a + get_formula_type result from this run. + """ + + accepted, rejected = _formula_attempts(ctx.messages) + if output.is_formula_valid: + if (output.table_id, output.field_name, output.formula) not in accepted: + raise ModelRetry( + f"{output.formula!r} was never accepted by " + f"{GET_FORMULA_TYPE_TOOL_NAME} for field {output.field_name!r} " + f"in table {output.table_id} in this run, so its validity is " + f"unverified. Call {GET_FORMULA_TYPE_TOOL_NAME} on it and return " + "the exact formula that passed." + ) + elif ( + len({_normalize_formula(formula) for _, _, formula in accepted} | rejected) + < FORMULA_MIN_ATTEMPTS_BEFORE_IMPOSSIBLE + ): + conversions = "; ".join( + f"to {target} use {how}" + for target, how in sorted(_CONVERSION_TO_TARGET_TYPE.items()) + ) + raise ModelRetry( + "A single attempt is not evidence that the request cannot be " + f"expressed. Validate at least {FORMULA_MIN_ATTEMPTS_BEFORE_IMPOSSIBLE} " + f"materially different candidates with {GET_FORMULA_TYPE_TOOL_NAME} " + "before giving up: vary the approach, and where a direct expression is " + "rejected try one that converts the argument types first — any field " + f"type can be converted ({conversions}). Justify failure only by " + f"quoting the error {GET_FORMULA_TYPE_TOOL_NAME} returned, never by " + "asserting from memory what the formula language does or does not " + "support." + ) + return output + + +# Keyed by the type name the compiler prints as the usable type for an argument. +_CONVERSION_TO_TARGET_TYPE: dict[str, str] = { + "text": "totext(x) for a single value, or join(x, ', ') for a list", + "char": "totext(x)", + "url": "tourl(totext(x))", + "link": "link(totext(x))", + "number": "tonumber(totext(x)), or count(x) to count a list", + "date": "todate(totext(x), 'YYYY-MM-DD')", + "duration": "toduration(tonumber(totext(x))) reading the number as seconds", + "boolean": "a comparison such as totext(x) != '' — there is no cast to boolean", +} + +_USABLE_TYPES = re.compile( + r"the only usable types? for this argument (?:is|are) ([a-z_]+(?:,[a-z_]+)*)" +) + + +def _type_mismatch_hint(error: str) -> str: + """ + Explains that a rejected argument type is a conversion problem. + + Without this the compiler's wording reads as the language not supporting the + operation at all, and the agent abandons a formula that a wrapped argument + would have made valid. + """ + + if "was of type" not in error: + return "" + + if "there are no possible types usable here" in error: + return ( + " That argument slot accepts no type at all, so no conversion will " + "fix it: restructure the expression instead of retrying conversions." + ) + + targets = {t for match in _USABLE_TYPES.findall(error) for t in match.split(",")} + repairs = [ + f"to {target} use {_CONVERSION_TO_TARGET_TYPE[target]}" + for target in sorted(targets) + if target in _CONVERSION_TO_TARGET_TYPE + ] + hint = ( + " This is an argument type mismatch, not an unsupported operation. Any " + "field type can be converted, so wrap the argument the error names in a " + "conversion function and validate again rather than abandoning the formula." + ) + if repairs: + hint += " Convert " + "; ".join(repairs) + "." + return hint + def get_formula_type_tool( user: AbstractUser, workspace: Workspace ) -> Callable[[str], str]: """ - Returns a function that validates a formula and returns its type. + Build the formula validation tool for the formula generation agent. + + :param user: The acting user, used to scope table access. + :param workspace: Workspace whose tables the formula may reference. + :returns: A function that validates a formula and returns its type. """ def get_formula_type(table_id: int, field_name: str, formula: str) -> str: @@ -87,21 +278,32 @@ def get_formula_type(table_id: int, field_name: str, formula: str) -> str: table = helpers.filter_tables(user, workspace).filter(id=table_id).first() if not table: - raise ValueError(f"Table with ID {table_id} not found in workspace.") + valid_ids = list( + helpers.filter_tables(user, workspace).values_list("id", flat=True) + ) + raise ModelRetry( + f"Table with ID {table_id} not found in workspace. " + f"Valid table IDs: {valid_ids}" + ) + # Only ModelRetry becomes a retry prompt; anything else aborts the turn. field = FormulaField(formula=formula, table=table, name=field_name, order=0) - field.recalculate_internal_fields(raise_if_invalid=True) - - result = TypeFormulaResultSerializer(field).data - if result["error"]: + try: + field.recalculate_internal_fields(raise_if_invalid=True) + result = TypeFormulaResultSerializer(field).data + error = result["error"] + except Exception as exc: + error = str(exc) + + if error: field_names = list( FieldHandler() .get_base_fields_queryset() .filter(table=table) .values_list("name", flat=True) ) - raise TypeError( - f"Invalid formula: {result['error']}. " + raise ModelRetry( + f"Invalid formula: {error}.{_type_mismatch_hint(error)} " f"Available fields in table '{table.name}': {', '.join(field_names)}" ) @@ -110,6 +312,42 @@ def get_formula_type(table_id: int, field_name: str, formula: str) -> str: return get_formula_type +def run_formula_generation( + user: AbstractUser, + workspace: Workspace, + prompt: str, + model_profile: "ResolvedAssistantModelProfile", +) -> AgentRunResult[FormulaGenerationResult]: + """ + Run the formula generation agent with its validation toolset and budgets. + + Failure policy is deliberately left to the caller: the fixer swallows + errors to None, while the generate_formula tool converts them to ModelRetry. + + :param user: The acting user, used to scope formula validation to their tables. + :param workspace: Workspace whose tables the formula may reference. + :param prompt: Fully formatted prompt for the agent. + :param model_profile: The model profile resolved for the assistant request. + :returns: The agent run result carrying a FormulaGenerationResult output. + """ + + from baserow_enterprise.assistant.model_profiles import UTILITY + + formula_toolset = FunctionToolset( + [Tool(get_formula_type_tool(user, workspace))], + max_retries=FORMULA_AGENT_RETRIES, + ) + model = model_profile.create_model() + return run_agent_sync_with_model( + formula_generation_agent, + prompt, + model=model, + model_settings=model_profile.get_settings(UTILITY), + toolsets=[formula_toolset], + usage_limits=UsageLimits(request_limit=FORMULA_AGENT_REQUEST_LIMIT), + ) + + def make_formula_fixer( user: AbstractUser, workspace: Workspace, tool_helpers ) -> Callable[[Any, str, str], str | None]: @@ -140,23 +378,17 @@ def fix_formula(table: Any, field_name: str, original_formula: str) -> str | Non _("Fixing formula for %(name)s...") % {"name": field_name} ) - formula_type_tool = Tool(get_formula_type_tool(user, workspace)) - formula_toolset = FunctionToolset([formula_type_tool]) prompt = format_formula_fixer_prompt( field_name, original_formula, schema, get_formula_docs() ) - from baserow_enterprise.assistant.model_profiles import UTILITY - - model_profile = tool_helpers.model_profile - model = model_profile.create_model() - result = run_agent_sync_with_model( - formula_generation_agent, - prompt, - model=model, - model_settings=model_profile.get_settings(UTILITY), - toolsets=[formula_toolset], - usage_limits=UsageLimits(request_limit=20), - ) + try: + result = run_formula_generation( + user, workspace, prompt, tool_helpers.model_profile + ) + except Exception: + # The fixer is best-effort and runs inside another except handler. + logger.exception("[assistant] formula fixer raised unexpectedly") + return None if result.output.is_formula_valid: return result.output.formula return None diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/prompts.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/prompts.py index cc1808ef37..9f6de474a2 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/prompts.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/prompts.py @@ -6,10 +6,39 @@ # Agent instructions # --------------------------------------------------------------------------- -FORMULA_AGENT_INSTRUCTIONS = ( - "Generates a Baserow formula based on the provided description and table schema. " - "Always validate the formula using the get_formula_type tool before returning it." -) +FORMULA_AGENT_INSTRUCTIONS = """\ +You write Baserow formulas. `get_formula_type` compiles a formula against the real +table and returns its type, or an error explaining what is wrong. It is the only +authority on what the formula language can and cannot do — your memory is not. + +For every field you are asked to produce: +1. Read the field types out of the schema in the prompt before writing anything. A + field's type decides which functions will accept it. +2. Write a candidate formula, using only functions listed in the Function Reference + of the formula documentation and obeying its Hard Rules. +3. When an argument's type is not one the function accepts, convert it rather than + abandoning the approach. `totext(x)` accepts a value of any type. `tonumber(x)` + accepts text, so `tonumber(totext(x))` reads a number out of a value of another + type. `join(x, ', ')` collapses a list of values — a link, lookup or other array + — into one string. `when_empty(x, fallback)` supplies a default. The Function + Reference lists what each function accepts. +4. Call `get_formula_type` on that candidate. Never return a formula you have not + validated in this run. +5. If it errors, the message names the problem and the fields available. Fix that + specific argument and validate again. + +Never state that Baserow "does not support" a conversion, a function or a field +type. If you believe a request cannot be expressed, prove it: at least one of your +validated attempts must have tried converting the argument types, and +`get_formula_type` must have rejected it. A belief that something is unsupported is +not a result. + +Only then set `is_formula_valid=false`, and put the verbatim text of the last +`get_formula_type` rejection into `error_message`, together with the formulas you +tried. Do not paraphrase it and do not invent a reason — that text is what the +caller sees. A validated formula that covers most of the request is better than +refusing outright. +""" SAMPLE_ROW_AGENT_INSTRUCTIONS = ( "Create 5 realistic sample rows for each table using the " diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/tools.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/tools.py index e0510d22dd..73b9f15dea 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/tools.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/tools.py @@ -6,15 +6,15 @@ from loguru import logger from pydantic import Field, create_model -from pydantic_ai import RunContext, Tool +from pydantic_ai import ModelRetry, RunContext, Tool from pydantic_ai.toolsets import FunctionToolset -from pydantic_ai.usage import UsageLimits from baserow.contrib.database.fields.actions import ( CreateFieldActionType, DeleteFieldActionType, UpdateFieldActionType, ) +from baserow.contrib.database.fields.handler import FieldHandler from baserow.contrib.database.fields.registries import field_type_registry from baserow.contrib.database.models import Database from baserow.contrib.database.rows.actions import ( @@ -29,20 +29,19 @@ UpdateViewFieldOptionsActionType, ) from baserow.contrib.database.views.handler import ViewHandler -from baserow.core.generative_ai.lifecycle import run_agent_sync_with_model from baserow.core.models import Workspace from baserow.core.service import CoreService from baserow_enterprise.assistant.deps import AssistantDeps +from baserow_enterprise.assistant.tools.shared import require_payload from baserow_enterprise.assistant.tools.toolset import inline_refs from baserow_enterprise.assistant.types import TableNavigationType, ViewNavigationType from baserow_premium.prompts import get_formula_docs from . import helpers from .agents import ( - formula_generation_agent, generate_sample_rows, - get_formula_type_tool, make_formula_fixer, + run_formula_generation, ) from .prompts import format_formula_generation_prompt from .types import ( @@ -476,14 +475,14 @@ def create_tables( DO NOT USE when: Tables already exist — check with list_tables first. HOW: Pass ALL related tables in a single call — link_row fields can reference other tables in the same call by name (they are created internally before fields are added). Choose appropriate field types for each column. Use single_select/multiple_select with select_options for categorical data. The primary field is always text — pick a meaningful name for it. + REQUIRED: `database_id` and `tables` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. """ user = ctx.deps.user workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not tables: - return {"created_tables": []} + require_payload("create_tables", "tables", tables) database = CoreService().get_application( user, @@ -499,6 +498,9 @@ def create_tables( user, tables, created_tables, tool_helpers, formula_fixer ) + # A link_row or lookup to an existing table adds a reverse field there. + _refresh_row_tools(ctx) + last_table = created_tables[-1] tool_helpers.navigate_to( TableNavigationType( @@ -569,14 +571,14 @@ def create_fields( RETURNS: Created fields with id, name, type. Formula errors with hints if any. DO NOT USE when: Creating a brand new table — use create_tables instead, which handles fields as part of table creation. HOW: Call get_tables_schema first to see existing fields and avoid duplicates. For link_row fields, ensure the target table already exists. + REQUIRED: `table_id` and `fields` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. """ user = ctx.deps.user workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not fields: - return {"created_fields": []} + require_payload("create_fields", "fields", fields) table = helpers.get_table(user, workspace, table_id) @@ -585,16 +587,19 @@ def create_fields( created_fields, field_errors, formula_errors = helpers.create_fields( user, table, fields, tool_helpers, formula_fixer=formula_fixer ) - result = {"created_fields": [field.model_dump() for field in created_fields]} - if field_errors: - result["field_errors"] = field_errors - if formula_errors: - for err in formula_errors: - err["hint"] = ( - "Use generate_formula to create a valid formula for this field." - ) - result["formula_errors"] = formula_errors - return result + + _refresh_row_tools(ctx) + + result = {"created_fields": [field.model_dump() for field in created_fields]} + if field_errors: + result["field_errors"] = field_errors + if formula_errors: + for err in formula_errors: + err["hint"] = ( + "Use generate_formula to create a valid formula for this field." + ) + result["formula_errors"] = formula_errors + return result # --------------------------------------------------------------------------- @@ -650,6 +655,8 @@ def update_fields( except Exception as e: errors.append(f"Error updating field {field_update.field_id}: {e}") + _refresh_row_tools(ctx) + result: dict[str, Any] = {"updated_fields": updated} if errors: result["errors"] = errors @@ -703,6 +710,8 @@ def delete_fields( except Exception as e: errors.append(f"Error deleting field {field_id}: {e}") + _refresh_row_tools(ctx) + result: dict[str, Any] = {"deleted_field_ids": deleted} if errors: result["errors"] = errors @@ -714,6 +723,43 @@ def delete_fields( # --------------------------------------------------------------------------- +def _describe_fields(table: Table) -> str: + """List a table's fields as ``id (name, type)`` for retry prompts.""" + + fields = FieldHandler().get_base_fields_queryset().filter(table=table) + return ( + ", ".join(f"{f.id} ({f.name}, {f.get_type().type})" for f in fields) or "none" + ) + + +def _drop_form_incompatible_fields( + table: Table, field_options: dict +) -> tuple[dict, list[str]]: + """ + Split form field_options into the ones the form view accepts and the names + it rejects. + + A form view refuses read-only fields and field types that cannot be a form + input (e.g. formula); enabling one raises and would roll the whole + create_views call back, so they are skipped and reported instead. + """ + + fields = {f.id: f for f in table.field_set.all()} + kept: dict = {} + skipped: list[str] = [] + for field_id, options in field_options.items(): + field = fields.get(int(field_id)) + if field is None: + kept[field_id] = options + continue + field_type = field_type_registry.get_by_model(field.specific_class) + if field.read_only or not field_type.can_be_in_form_view: + skipped.append(field.name) + else: + kept[field_id] = options + return kept, skipped + + def create_views( ctx: RunContext[AssistantDeps], table_id: Annotated[ @@ -737,14 +783,14 @@ def create_views( RETURNS: Created views with id, name, type configuration. DO NOT USE when: The default grid view already meets the user's needs. Check existing views with list_views to avoid duplicates. HOW: Each view type requires specific config. Form views: provide field_options listing every field to show (field_id, name, order, required). Kanban: set column_field_id to a single_select field. Calendar: set date_field_id to a date field. Timeline: set both start/end date fields. Gallery: optionally set cover_field_id to a file field. Call get_tables_schema first to get the field IDs you need. + REQUIRED: `table_id` and `views` must arrive in the same call. The ID says where to act; the payload says what to create. A call carrying only the ID creates nothing and is rejected. """ user = ctx.deps.user workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not views: - return {"created_views": []} + require_payload("create_views", "views", views) table = helpers.get_table(user, workspace, table_id) @@ -757,18 +803,51 @@ def create_views( % {"view_type": view.type, "view_name": view.name} ) + # A wrong field id must become a retry prompt, not a turn-ending error. + try: + orm_kwargs = view.to_django_orm_kwargs(table) + except (ValueError, TypeError) as exc: + raise ModelRetry( + f"Cannot create the '{view.name}' {view.type} view: {exc} " + f"Fields in table '{table.name}': {_describe_fields(table)}" + ) from exc + orm_view = CreateViewActionType.do( user, table, view.type, - **view.to_django_orm_kwargs(table), + **orm_kwargs, ) field_options = view.field_options_to_django_orm() + skipped_fields: list[str] = [] if field_options: - UpdateViewFieldOptionsActionType.do(user, orm_view, field_options) - - created_views.append({"id": orm_view.id, **view.model_dump()}) + field_options, skipped_fields = _drop_form_incompatible_fields( + table, field_options + ) + if field_options: + try: + UpdateViewFieldOptionsActionType.do(user, orm_view, field_options) + except Exception as exc: + # ModelRetry rolls the transaction back, so say so explicitly. + raise ModelRetry( + f"The field_options of the '{view.name}' {view.type} view " + f"were rejected and no views were created: {exc} Retry " + f"without the rejected fields. Fields in table " + f"'{table.name}': {_describe_fields(table)}" + ) from exc + + created = {"id": orm_view.id, **view.model_dump()} + if view.type == "form": + created["field_options"] = ViewItem.from_django_orm( + orm_view + ).model_dump()["field_options"] + if skipped_fields: + created["skipped_fields"] = ( + "Not shown on the form, these field types cannot be a form " + f"input: {', '.join(skipped_fields)}" + ) + created_views.append(created) tool_helpers.navigate_to( ViewNavigationType( @@ -809,6 +888,7 @@ def create_view_filters( RETURNS: Created filters with id and configuration per view. DO NOT USE when: The view doesn't exist yet — create it first with create_views. HOW: Get the table schema first to know field IDs and types. Match filter type to field type. + REQUIRED: `view_filters` must contain at least one entry, each carrying the ID of the view it applies to. The ID says where to act; the payload says what to create. A call carrying no entries creates nothing and is rejected. ## Value formats by type @@ -824,8 +904,7 @@ def create_view_filters( workspace = ctx.deps.workspace tool_helpers = ctx.deps.tool_helpers - if not view_filters: - return {"created_view_filters": []} + require_payload("create_view_filters", "view_filters", view_filters) created_view_filters = [] for vf in view_filters: @@ -897,9 +976,9 @@ def generate_formula( :param save_to_field: Whether to save the formula to a field. :param thought: Brief reasoning for invoking the tool. :return: Formula metadata, including an integer table ID when a field is saved. - :raises Exception: If the formula is invalid or targets an unavailable table. + :raises ModelRetry: If generation fails, the formula is invalid, or its target + table is unavailable. """ - from baserow_enterprise.assistant.model_profiles import UTILITY user = ctx.deps.user workspace = ctx.deps.workspace @@ -915,33 +994,41 @@ def generate_formula( tool_helpers.update_status(_("Generating formula...")) formula_docs = get_formula_docs() - formula_type_tool = Tool(get_formula_type_tool(user, workspace)) - formula_toolset = FunctionToolset([formula_type_tool]) - prompt = format_formula_generation_prompt( description, database_tables_schema, formula_docs ) - model_profile = tool_helpers.model_profile - model = model_profile.create_model() - agent_result = run_agent_sync_with_model( - formula_generation_agent, - prompt, - model=model, - model_settings=model_profile.get_settings(UTILITY), - toolsets=[formula_toolset], - usage_limits=UsageLimits(request_limit=20), - ) + try: + agent_result = run_formula_generation( + user, workspace, prompt, tool_helpers.model_profile + ) + except ModelRetry: + # A retry prompt for the model, not a failure to rewrite as one below. + raise + except Exception as exc: + # A sub-agent failure must not end the turn; create_tables does the same. + logger.exception("[assistant] formula_generation_agent raised unexpectedly") + raise ModelRetry( + f"The formula generator failed: {exc}. Retry with a simpler " + "description, or create the field without a formula." + ) from exc + result = agent_result.output + # Recoverable: the orchestrator can rephrase or retarget; raising ends the turn. if not result.is_formula_valid: - raise Exception(f"Error generating formula: {result.error_message}") + raise ModelRetry( + f"Could not generate a valid formula: {result.error_message} " + "Rephrase the description, or tell the user which part is not " + "expressible in the Baserow formula language." + ) table = next((t for t in database_tables if t.id == result.table_id), None) if table is None: - raise Exception( - "The generated formula is intended for a different table " - f"than the current one. Table with ID {result.table_id} not found." + valid = ", ".join(f"{t.id} ({t.name})" for t in database_tables) + raise ModelRetry( + f"The generated formula targets table {result.table_id}, which is not " + f"in database {database_id}. Tables available: {valid}" ) data: dict[str, str | int] = { @@ -998,6 +1085,9 @@ def generate_formula( } ) + # Saving over a non-formula field trashes it, changing the row schema. + _refresh_row_tools(ctx) + return data @@ -1038,8 +1128,7 @@ def _create_rows( ) -> dict[str, Any]: """Create new rows in the specified table.""" - if not rows: - return {"created_row_ids": []} + require_payload(f"create_rows_in_table_{table.id}", "rows", rows) tool_helpers.update_status( _("Creating rows in %(table_name)s ") % {"table_name": table.name} @@ -1055,13 +1144,19 @@ def _create_rows( create_rows_tool = Tool( _create_rows, name=f"create_rows_in_table_{table.id}", + metadata={"table_id": table.id}, description=( f"WHEN: Creating new rows in '{table.name}' (ID: {table.id}). " - f"WHAT: Inserts up to 20 rows with field values matching the table schema. " + f"WHAT: Inserts rows with field values matching the table schema. " + f"Keep batches to at most 20 rows to avoid oversized generated arguments; " + f"for more rows, call this tool again with the next batch. " f"RETURNS: Created row IDs. " f"DO NOT USE: For other tables — each table has its own create tool. " f"HOW: Fill EVERY field including ALL link_row (relationship) fields. Never skip a field unless data is genuinely unavailable." f"{link_row_hints}" + f" REQUIRED: `rows` must contain at least one row. This tool already " + f"knows where to write, so the call must also carry what to create; a " + f"call with an empty `rows` creates nothing and is rejected." ), max_retries=2, ) @@ -1092,6 +1187,7 @@ def _update_rows( update_rows_tool = Tool( _update_rows, name=f"update_rows_in_table_{table.id}", + metadata={"table_id": table.id}, description=( f"WHEN: Updating existing rows in '{table.name}' (ID: {table.id}) by row ID. " f"WHAT: Updates specified fields on up to 20 rows. Only include fields you want to change — omit fields to keep them unchanged. " @@ -1126,6 +1222,7 @@ def _delete_rows( delete_rows_tool = Tool( _delete_rows, name=f"delete_rows_in_table_{table.id}", + metadata={"table_id": table.id}, description=( f"WHEN: Deleting rows from '{table.name}' (ID: {table.id}) by row ID. " f"WHAT: Permanently removes up to 20 specified rows. " @@ -1141,6 +1238,51 @@ def _delete_rows( } +def _refresh_row_tools(ctx: RunContext[AssistantDeps]) -> None: + """ + Rebuild the loaded row tools so their schema matches the table's fields. + + Row tools bake the table schema into their signature when loaded, so a field + change leaves them rejecting new fields or writing to dropped columns. Every + loaded table is rebuilt, because a link_row field also adds a reverse field + to the table it points at. + + :param ctx: The run context holding the dynamic tool registry. + """ + + dynamic_tools = ctx.deps.dynamic_tools + if not dynamic_tools: + return + + table_ids = { + tool.metadata["table_id"] + for tool in dynamic_tools + if tool.metadata and "table_id" in tool.metadata + } + tables = helpers.filter_tables(ctx.deps.user, ctx.deps.workspace).filter( + id__in=table_ids + ) + + rebuilt: dict[str, Tool] = {} + for table in tables: + try: + row_tools = _build_row_tools( + ctx.deps.user, ctx.deps.workspace, ctx.deps.tool_helpers, table + ) + except Exception: + # Raising would fail a field change that already succeeded. + logger.exception( + "[assistant] could not refresh row tools for table {}", table.id + ) + continue + # delete_rows takes row IDs only, so its signature cannot go stale. + rebuilt.update( + {row_tools[op].name: row_tools[op] for op in ("create", "update")} + ) + + dynamic_tools[:] = [rebuilt.get(tool.name, tool) for tool in dynamic_tools] + + # --------------------------------------------------------------------------- # Tool 10: load_row_tools # --------------------------------------------------------------------------- @@ -1167,7 +1309,8 @@ def load_row_tools( WHEN to use: You need to directly create, update, or delete rows in a database table. Must be called before any row manipulation. WHAT it does: Unlocks table-specific tools and their schema: create_rows_in_table_X, update_rows_in_table_X, delete_rows_in_table_X for each table ID provided. The loaded tools include the full field schema — no need to call get_tables_schema. RETURNS: Names of newly available tools. - DO NOT USE when: Row tools for these tables are already loaded from a previous call in this session. + DO NOT USE when: Row tools for these tables are already loaded and no field has changed since — one call per table is enough. + AFTER A SCHEMA CHANGE: create_fields, update_fields and delete_fields refresh the loaded row tools automatically, so no reload is needed. If a row tool still rejects a field you just created, call this again: reloading replaces the stale tools with ones matching the current schema. DO NOT USE for builder workflow actions — if you want a button/form in an Application Builder page to create/update/delete rows, use create_actions instead. load_row_tools is for direct database manipulation, NOT for configuring app behavior. HOW: Just call this with the table ID(s) and operations you need. The loaded row tools already contain the complete field schema in their parameters — do NOT call get_tables_schema or search_user_docs before or after this tool. @@ -1200,9 +1343,11 @@ def load_row_tools( if "delete" in operations: new_tools.append(table_tools["delete"]) - # Store new tools in dynamic_tools for the dynamic toolset - # to pick up on the next agent step - ctx.deps.dynamic_tools.extend(new_tools) + # Replaced by name: reloads must see schema changes, and add_tool rejects dupes. + new_names = {t.name for t in new_tools} + ctx.deps.dynamic_tools[:] = [ + t for t in ctx.deps.dynamic_tools if t.name not in new_names + ] + new_tools tool_names = [t.name for t in new_tools] return f"Tools loaded: {', '.join(tool_names)}" diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/types/views.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/types/views.py index fdbf2d7fdd..25f2e72639 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/database/types/views.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/database/types/views.py @@ -84,7 +84,7 @@ class GridFieldOption(BaseModel): def _grid_to_orm(v, table): - return {"row_height": v.row_height} + return {"row_height_size": v.row_height} def _kanban_to_orm(v, table): @@ -108,7 +108,7 @@ def _gallery_to_orm(v, table): cover_field = model.get_field_object_by_id(v.cover_field_id)["field"] if not isinstance(cover_field, FileField): raise ValueError("The cover_field_id must be a File field.") - return {"card_cover_image_field_id": v.cover_field_id} + return {"card_cover_image_field": cover_field} def _timeline_to_orm(v, table): diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/registries.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/registries.py index 8c6b374708..9a3d112933 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/registries.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/registries.py @@ -134,6 +134,11 @@ def build_toolset( AgentMode.AUTOMATION, routing_rules_by_type.get("automation", ""), ), + ( + "explain", + AgentMode.EXPLAIN, + routing_rules_by_type.get("search_user_docs", ""), + ), ] manifests = {} @@ -145,28 +150,28 @@ def build_toolset( ] manifest = generate_tool_manifest_compact(groups, routing_rules=rules) - # Append a compact cross-mode summary so the agent knows what - # capabilities exist in other modes (and can switch_mode to use them). + # Each gated tool is listed under a callable mode, never silently absent. + listed: set[str] = set() other_lines = [] for other_key, other_mode, _ in _mode_config: if other_key == mode_key: continue - specific = mode_map[other_mode] - shared - other_lines.append(f"- {other_key}: {', '.join(sorted(specific))}") + elsewhere = sorted(mode_map[other_mode] - shared - allowed - listed) + if not elsewhere: + continue + listed.update(elsewhere) + other_lines.append( + f'- switch_mode("{other_key}") to call: {", ".join(elsewhere)}' + ) if other_lines: - manifest += "\n\n## Other modes (switch_mode to access)\n" + "\n".join( - other_lines + manifest += ( + "\n\n## Tools that exist but are not callable in this mode\n" + "These tools are part of Baserow. To use one, call switch_mode " + "for its mode first, then call the tool.\n" + "\n".join(other_lines) ) manifests[mode_key] = manifest - explain_allowed = mode_map[AgentMode.EXPLAIN] - explain_groups = [ - (label, [f for f in funcs if f.__name__ in explain_allowed]) - for label, funcs in module_groups - ] - manifests["explain"] = generate_tool_manifest_compact(explain_groups) - return ( InlineRefsToolset(mode_aware, model=model, model_profile=model_profile), manifests["database"], diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/__init__.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/__init__.py index de86534aa3..96c36b84b0 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/__init__.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/__init__.py @@ -15,6 +15,7 @@ needs_formula, wrap_static_string, ) +from .payloads import require_payload __all__ = [ "EMPTY_FORMULA", @@ -32,4 +33,5 @@ "create_example_from_json_schema", "BaseFormulaContext", "get_formula_generator", + "require_payload", ] diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/payloads.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/payloads.py new file mode 100644 index 0000000000..f401f61a54 --- /dev/null +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/shared/payloads.py @@ -0,0 +1,27 @@ +"""Payload guard shared by every collection-creating tool.""" + +from typing import Any + +from pydantic_ai import ModelRetry + + +def require_payload(tool_name: str, arg_name: str, items: Any) -> None: + """Reject a call that named a target but carried nothing to act on. + + An empty success result would let a dropped payload read as done work. + + :param tool_name: The collection-creating tool receiving the call. + :param arg_name: The argument containing the items to create. + :param items: The collection supplied for that argument. + :return: None when the collection is nonempty. + :raises ModelRetry: When the payload is empty and must be supplied again. + """ + + if not items: + raise ModelRetry( + f"{tool_name} received an empty `{arg_name}`. `{arg_name}` is " + f"required and must contain at least one item: an ID argument " + f"only says where to act, it never says what to create. Nothing " + f"was changed. Resend the call with the target ID and the full " + f"`{arg_name}` list together." + ) diff --git a/enterprise/backend/src/baserow_enterprise/assistant/tools/toolset.py b/enterprise/backend/src/baserow_enterprise/assistant/tools/toolset.py index 4ef3181696..4d5789a8a2 100644 --- a/enterprise/backend/src/baserow_enterprise/assistant/tools/toolset.py +++ b/enterprise/backend/src/baserow_enterprise/assistant/tools/toolset.py @@ -79,16 +79,236 @@ def _resolve(node, *, _inside_properties=False): return _resolve(schema) +# --------------------------------------------------------------------------- +# Validation-error rendering +# --------------------------------------------------------------------------- + +# Read-only tool that returns real values for a given id argument. +_ID_DISCOVERY_TOOL: dict[str, str] = { + "table_id": "list_tables", + "field_id": "get_tables_schema", + "column_field_id": "get_tables_schema", + "date_field_id": "get_tables_schema", + "cover_field_id": "get_tables_schema", + "start_date_field_id": "get_tables_schema", + "end_date_field_id": "get_tables_schema", + "view_id": "list_views", + "row_id": "list_rows", + "database_id": "list_builders", + "application_id": "list_builders", + "automation_id": "list_builders", + "builder_id": "list_builders", + "page_id": "list_pages", + "navigate_to_page_id": "list_pages", + "element_id": "list_elements", + "data_source_id": "list_data_sources", + "workflow_id": "list_workflows", + "node_id": "list_nodes", +} + +_MAX_REPORTED_ERRORS = 8 + + +def _unwrap_union(node: Any) -> dict: + """First non-null branch of an anyOf/oneOf, else the node itself.""" + + if not isinstance(node, dict): + return {} + for key in ("anyOf", "oneOf"): + for branch in node.get(key) or (): + if isinstance(branch, dict) and branch.get("type") != "null": + return branch + return node + + +def _schema_at(schema: dict, loc: tuple) -> tuple[dict, bool]: + """Follow a Pydantic error location into a schema with inlined references. + + Union branch tags do not resolve, so a partial result must not be used + as the authoritative key set. + + :param schema: The tool's parameter schema after reference inlining. + :param loc: Property names and array indices from a validation error. + :return: The deepest node reached and whether the full location resolved. + """ + + node = _unwrap_union(schema) + for part in loc: + nxt = ( + _unwrap_union(node.get("items")) + if isinstance(part, int) + else _unwrap_union((node.get("properties") or {}).get(part)) + ) + if not nxt: + return node, False + node = nxt + return node, True + + +def _keys_of(node: dict) -> list[str]: + """Property names of *node*, required ones suffixed with ``*``.""" + + required = set(node.get("required") or ()) + return [f"{k}*" if k in required else k for k in (node.get("properties") or {})] + + +def _describe_shape(node: dict, exact: bool) -> str: + """Describe an expected value without guessing from a partial schema match. + + :param node: The schema node reached while resolving an error location. + :param exact: Whether the complete error location resolved to this node. + :return: A brief expected shape, or a generic schema reference. + """ + + if not node or not exact: + return "the value this tool's schema documents at that path" + if node.get("type") == "array": + item = _unwrap_union(node.get("items")) + keys = _keys_of(item) + if keys: + return f"a list of objects with keys: {', '.join(keys)}" + return f"a list of {item.get('type') or 'values'}" + if node.get("enum"): + return "one of: " + ", ".join(str(v) for v in node["enum"]) + keys = _keys_of(node) + if keys: + return f"an object with keys: {', '.join(keys)}" + return str(node.get("type") or "value") + + +def _discovery_hint(field_name: str) -> str: + """Suggest an ID lookup tool, or return no hint for other fields.""" + + tool = _ID_DISCOVERY_TOOL.get(field_name) + if tool: + return f" Call {tool} to get a real id." + if field_name.endswith("_id"): + return " Call the matching list_* tool to get a real id." + return "" + + +def _short(value: Any, limit: int = 60) -> str: + """Render a value for error feedback, truncating it after ``limit`` characters.""" + + try: + text = json.dumps(value, default=str) + except (TypeError, ValueError): + text = repr(value) + return text if len(text) <= limit else f"{text[:limit]}…(truncated)" + + +def format_tool_arg_errors( + tool_name: str, schema: dict, wrong_args: Any, errors: list[dict] +) -> str: + """Render Pydantic errors as recovery instructions, with a fallback on failure. + + :param tool_name: The tool whose arguments failed validation. + :param schema: The tool's parameter schema after reference inlining. + :param wrong_args: The raw arguments that failed validation. + :param errors: Error details returned by Pydantic validation. + :return: Correction instructions, or the original errors if rendering fails. + """ + + try: + return _render_tool_arg_errors(tool_name, schema, wrong_args, errors) + except Exception: + logger.exception("[assistant] Could not render arg errors for '{}'", tool_name) + return f"{tool_name} did NOT run — its arguments were rejected: {errors}" + + +def _render_tool_arg_errors( + tool_name: str, schema: dict, wrong_args: Any, errors: list[dict] +) -> str: + """Build a bounded error report with expected shapes and ID lookup hints. + + :param tool_name: The tool to call again with corrected arguments. + :param schema: The tool's parameter schema after reference inlining. + :param wrong_args: The rejected arguments, used to report the supplied keys. + :param errors: Pydantic errors, capped at ``_MAX_REPORTED_ERRORS`` in the report. + :return: A report describing rejected values and how to correct them. + """ + + lines: list[str] = [] + id_rejected = False + + for err in errors[:_MAX_REPORTED_ERRORS]: + loc = tuple(err.get("loc") or ()) + path = ".".join(str(p) for p in loc) or "(arguments)" + leaf = str(loc[-1]) if loc else "" + id_rejected = id_rejected or leaf.endswith("_id") + err_type = err.get("type", "") + + if err_type == "missing": + node, exact = _schema_at(schema, loc) + lines.append( + f"- {path}: required, but you did not send it. Send " + f"{_describe_shape(node, exact)}.{_discovery_hint(leaf)}" + ) + elif err_type in ("extra_forbidden", "unexpected_keyword_argument"): + node, exact = _schema_at(schema, loc[:-1]) + keys = _keys_of(node) if exact else [] + allowed = ( + f"The only keys accepted here are: {', '.join(keys)}." + if keys + else "Use only the keys this tool's schema defines at that path." + ) + lines.append( + f"- {path}: '{leaf}' is not a key of this object. {allowed} " + "Move the value under an accepted key, or drop it." + ) + else: + node, exact = _schema_at(schema, loc) + lines.append( + f"- {path}: {err.get('msg', err_type)} — you sent " + f"{_short(err.get('input'))}, expected " + f"{_describe_shape(node, exact)}.{_discovery_hint(leaf)}" + ) + + hidden = len(errors) - len(lines) + if hidden > 0: + lines.append(f"- ...and {hidden} more error(s); fix them the same way.") + + sent = ( + ", ".join(sorted(wrong_args)) + if isinstance(wrong_args, dict) and wrong_args + else "(none)" + ) + footer = ( + "Keys marked * are required. Send the whole corrected argument object " + f"in a new {tool_name} call." + ) + if id_rejected: + footer += ( + " Never invent an id or send a placeholder such as 0 — take ids from " + "a create_* result or a list_* tool result." + ) + return ( + f"{tool_name} did NOT run — its arguments were rejected.\n" + f"Keys you sent: {sent}.\n" + "\n".join(lines) + f"\n{footer}" + ) + + # --------------------------------------------------------------------------- # Lenient validator & fixer # --------------------------------------------------------------------------- _FIXER_PROMPT = """\ -You are a JSON repair tool. You receive a JSON object that failed schema \ -validation, the validation errors, and the target JSON schema. Return ONLY \ -the fixed JSON object — no explanation, no markdown fences. Preserve the \ -original values as much as possible; only change what is needed to satisfy \ -the schema.""" +You repair tool-call JSON. You receive the target JSON schema, the object that \ +failed validation, and the validation errors. Return ONLY a JSON object — no \ +explanation, no markdown fences. + +Rules: +1. Preserve every value the caller supplied. You may move a value to the \ +correct key, rename a key, drop an unsupported key, or convert a value to the \ +type the schema requires — nothing else. +2. Do not invent data. Fill in a required value only when it is already \ +present elsewhere in the caller's object, or when the schema itself \ +determines it (a default, a const, or a single-member enum). Never invent an \ +identifier, and never satisfy a required field with a placeholder such as 0, \ +1, "", "unknown" or a guessed name. +3. If a required value is genuinely absent and rule 2 does not supply it, do \ +not guess. Return exactly \ +{"__cannot_fix__": ""}.""" class _LenientValidator: @@ -112,6 +332,87 @@ def validate_python(self, input, *, allow_partial="off", **kwargs): _LENIENT_VALIDATOR = _LenientValidator() +# --------------------------------------------------------------------------- +# Invented resource ids +# --------------------------------------------------------------------------- + +# Only ids listed here are checked, so a user's own ``*_id`` column is never flagged. +_ID_PRODUCERS: dict[str, str] = { + "database_id": "list_builders", + "application_id": "list_builders", + "automation_id": "list_builders", + "builder_id": "list_builders", + "table_id": "list_tables", + "field_id": "get_tables_schema", + "view_id": "list_views", + "page_id": "list_pages", + "element_id": "list_elements", + "data_source_id": "list_data_sources", + "workflow_id": "list_workflows", + "node_id": "list_nodes", +} + + +def _id_producer(key: str) -> str | None: + """Longest-suffix match, so ``cover_field_id`` resolves like ``field_id``.""" + + for name in sorted(_ID_PRODUCERS, key=len, reverse=True): + if key.endswith(name): + return _ID_PRODUCERS[name] + return None + + +def _is_placeholder_id(value: Any) -> bool: + """Identify non-positive IDs, leaving malformed values for schema validation. + + :param value: A raw tool argument that refers to a Baserow resource. + :return: Whether the argument represents an integer ID at or below zero. + """ + + if value is None or isinstance(value, bool): + return False + if isinstance(value, int): + return value <= 0 + if isinstance(value, str): + try: + return int(value) <= 0 + except ValueError: + return False + return False + + +def _find_placeholder_ids(node: Any, path: str = "") -> list[tuple[str, str, Any]]: + """Find non-positive resource IDs in nested raw tool arguments. + + :param node: The argument value to inspect recursively. + :param path: The JSON path leading to this value, empty at the root. + :return: ``(json_path, producer_tool, value)`` tuples for rejected resource IDs. + """ + + found: list[tuple[str, str, Any]] = [] + if isinstance(node, dict): + for key, value in node.items(): + child = f"{path}.{key}" if path else key + producer = _id_producer(key) + if producer is not None: + if _is_placeholder_id(value): + found.append((child, producer, value)) + continue + plural = _id_producer(key[:-1]) if key.endswith("s") else None + if plural is not None and isinstance(value, list): + found.extend( + (f"{child}[{i}]", plural, item) + for i, item in enumerate(value) + if _is_placeholder_id(item) + ) + continue + found.extend(_find_placeholder_ids(value, child)) + elif isinstance(node, list): + for i, item in enumerate(node): + found.extend(_find_placeholder_ids(item, f"{path}[{i}]")) + return found + + # --------------------------------------------------------------------------- # InlineRefsToolset # --------------------------------------------------------------------------- @@ -178,11 +479,8 @@ async def get_tools(self, ctx) -> dict[str, ToolsetTool[AgentDepsT]]: tool.tool_def.parameters_json_schema = inline_refs( tool.tool_def.parameters_json_schema ) - # Save the original validator and schema once, then replace with - # lenient passthrough so validation failures reach call_tool() - # where we can attempt an async fix. Guard against multiple calls - # so we don't overwrite the real validator with _LENIENT_VALIDATOR. - if name not in self._original_validators: + # Re-cache on every re-issue so the fixer sees the schema the model saw. + if tool.args_validator is not _LENIENT_VALIDATOR: self._original_validators[name] = tool.args_validator self._schemas[name] = tool.tool_def.parameters_json_schema tool.args_validator = _LENIENT_VALIDATOR @@ -195,6 +493,27 @@ async def call_tool( ctx: Any, tool: ToolsetTool[AgentDepsT], ) -> Any: + placeholders = _find_placeholder_ids(tool_args) + if placeholders: + logger.warning( + "[assistant] Tool '{}' called with placeholder IDs: {}", + name, + placeholders, + ) + offenders = ", ".join(f"{p}={v!r}" for p, _, v in placeholders) + producers = " / ".join(sorted({t for _, t, _ in placeholders})) + return { + "error": ( + f"Not executed. Invented IDs: {offenders}. Baserow IDs start " + "at 1. Only send an ID you have read from a tool result or " + "from ." + ), + "next_steps": ( + f"Call {producers} to read the real ID (switch_mode first if " + f"it is not available in this mode), then call {name} again " + "with the same arguments and that ID." + ), + } original_validator = self._original_validators.get(name) if original_validator: try: @@ -217,6 +536,7 @@ async def _fix_tool_args( schema = self._schemas.get(tool_name, {}) error_details = error.errors(include_url=False, include_context=False) + report = format_tool_arg_errors(tool_name, schema, wrong_args, error_details) logger.warning( "[assistant] Tool '{}' args failed validation, attempting fix. Errors: {}", @@ -228,7 +548,8 @@ async def _fix_tool_args( f"Tool: {tool_name}\n\n" f"Schema:\n{json.dumps(schema, indent=2)}\n\n" f"Invalid input:\n{json.dumps(wrong_args, indent=2)}\n\n" - f"Validation errors:\n{json.dumps(error_details, indent=2)}" + f"Validation errors:\n{json.dumps(error_details, indent=2)}\n\n" + f"What is wrong:\n{report}" ) try: @@ -256,23 +577,35 @@ async def _fix_tool_args( tool_name, exc, ) + raise ModelRetry(report) from exc + + if isinstance(fixed_args, dict) and "__cannot_fix__" in fixed_args: + missing = fixed_args["__cannot_fix__"] + logger.warning( + "[assistant] Fixer could not repair args for tool '{}': {}", + tool_name, + missing, + ) raise ModelRetry( - f"Tool arguments invalid and fix attempt failed: {error_details}" - ) from exc + f"Tool '{tool_name}' was called without required information: " + f"{missing}. Do not retry with a guessed or placeholder value — " + f"obtain the real value first (list or fetch the resource, or " + f"ask the user), then call the tool again." + ) # Re-validate with original schema original_validator = self._original_validators[tool_name] try: validated = original_validator.validate_python(fixed_args) except ValidationError as e2: + retry_errors = e2.errors(include_url=False, include_context=False) logger.warning( "[assistant] Fixed args for tool '{}' still invalid: {}", tool_name, - e2.errors(include_url=False, include_context=False), + retry_errors, ) raise ModelRetry( - f"Tool arguments still invalid after fix attempt: " - f"{e2.errors(include_url=False, include_context=False)}" + format_tool_arg_errors(tool_name, schema, fixed_args, retry_errors) ) from e2 return validated @@ -394,6 +727,28 @@ async def call_tool( "the correct IDs and retry." ) } + except Exception as exc: + # Every Baserow lookup miss raises a plain `DoesNotExist`. + if not type(exc).__name__.endswith("DoesNotExist"): + raise + logger.warning( + "[assistant] Tool '{}' referenced a missing resource: {}: {}", + name, + type(exc).__name__, + exc, + ) + return { + "error": ( + f"{name} referenced something that does not exist or is not " + f"accessible: {exc}" + ), + "next_steps": ( + "Every id and type value must come from a previous tool " + "result, never from a guess or a placeholder. Re-read it " + "from the latest list_*/get_*/create_* result, or call the " + f"matching list_*/get_* tool to look it up, then retry {name}." + ), + } # --------------------------------------------------------------------------- diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant.py index 675961d95f..2c597c2901 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant.py @@ -11,6 +11,8 @@ from asgiref.sync import async_to_sync from pydantic_ai.messages import PartStartEvent from pydantic_ai.messages import TextPart as PaiTextPart +from pydantic_ai.models.function import FunctionModel +from pydantic_ai.toolsets import FunctionToolset from baserow.core.ai_provider.constants import ( AI_PROVIDER_FEATURE_KUMA, @@ -37,7 +39,11 @@ get_model_string, resolve_assistant_model, ) -from baserow_enterprise.assistant.models import AssistantChat, AssistantChatMessage +from baserow_enterprise.assistant.models import ( + AssistantChat, + AssistantChatMessage, + AssistantChatPrediction, +) from baserow_enterprise.assistant.prompts import AGENT_SYSTEM_PROMPT from baserow_enterprise.assistant.types import ( AiMessage, @@ -54,6 +60,8 @@ WorkspaceUIContext, ) +from .utils import make_test_ctx + @pytest.fixture(autouse=True) def mock_posthog(): @@ -278,6 +286,28 @@ def test_load_message_history_returns_none_for_empty(self, enterprise_data_fixtu history = async_to_sync(assistant._load_message_history)() assert history is None + def test_save_ai_response_persists_posthog_trace_id(self, enterprise_data_fixture): + user = enterprise_data_fixture.create_user() + workspace = enterprise_data_fixture.create_workspace(user=user) + chat = AssistantChat.objects.create( + user=user, workspace=workspace, title="Test Chat" + ) + human_message = AssistantChatMessage.objects.create( + chat=chat, + role=AssistantChatMessage.Role.HUMAN, + content="Create a table", + ) + assistant = Assistant(chat) + assistant._telemetry.trace_id = "trace-123" + + async_to_sync(assistant._save_ai_response)(human_message, "Done") + + prediction = AssistantChatPrediction.objects.get(human_message=human_message) + assert prediction.prediction == { + "answer": "Done", + "posthog_trace_id": "trace-123", + } + def test_load_message_history_deserializes_and_compacts( self, enterprise_data_fixture ): @@ -476,6 +506,22 @@ def test_agent_system_prompt_includes_grounding_guardrail(self): assert "Use `search_user_docs` first" in AGENT_SYSTEM_PROMPT assert "Never invent plan names" in AGENT_SYSTEM_PROMPT + def test_agent_system_prompt_covers_production_regressions(self): + assert AGENT_SYSTEM_PROMPT.index("") < AGENT_SYSTEM_PROMPT.index( + "" + ) + assert "call create_builders first and build on the ID it returns" in ( + AGENT_SYSTEM_PROMPT + ) + assert "Never invent, guess, or carry over an ID from a different resource" in ( + AGENT_SYSTEM_PROMPT + ) + assert "Baserow IDs start at 1, so 0 is never an ID" in AGENT_SYSTEM_PROMPT + assert "For database formula creation or repair, call generate_formula" in ( + AGENT_SYSTEM_PROMPT + ) + assert "Never return or save a handwritten formula" in AGENT_SYSTEM_PROMPT + @pytest.mark.django_db class TestGetWorkspaceLicenseType: @@ -1329,3 +1375,72 @@ def test_database_model_readiness_is_cached_per_configuration( assert create_model.call_count == 2 assert test_model.call_count == 2 + + +@pytest.mark.asyncio +class TestAssistantTextToolCallRecovery: + @pytest.fixture + def assistant_for_responses(self): + def build(*responses): + calls = [] + + async def stream(messages, info): + calls.append(messages) + yield responses[min(len(calls) - 1, len(responses) - 1)] + + assistant = Assistant.__new__(Assistant) + assistant._model = FunctionModel(stream_function=stream) + assistant._model_profile = MagicMock() + assistant._model_profile.get_settings.return_value = {} + assistant._toolset = FunctionToolset() + assistant._deps = make_test_ctx(None, None).deps + return assistant, calls + + return build + + @pytest.mark.parametrize("fenced", [False, True]) + async def test_repeated_text_tool_calls_return_the_graceful_fallback( + self, assistant_for_responses, fenced + ): + payload = '{"name": "create_tables", "arguments": {"tables": []}}' + if fenced: + payload = f"```json\n{payload}\n```" + assistant, calls = assistant_for_responses(payload) + + answer, _result = await assistant._run_agent_with_retries( + "Create a table", None, asyncio.Queue() + ) + + assert answer == ( + "I ran into a temporary issue processing " + "your request. Could you please try again?" + ) + assert len(calls) == 3 + + async def test_a_corrected_answer_is_accepted_on_the_next_pass( + self, assistant_for_responses + ): + assistant, calls = assistant_for_responses( + '```json\n{"name": "list_tables", "arguments": {}}\n```', + "Which database should I use?", + ) + + answer, _result = await assistant._run_agent_with_retries( + "List the tables", None, asyncio.Queue() + ) + + assert answer == "Which database should I use?" + assert len(calls) == 2 + + async def test_ordinary_answers_are_accepted_without_retrying( + self, assistant_for_responses + ): + expected = 'The field {"name": "Arguments"} maps to your schema.' + assistant, calls = assistant_for_responses(expected) + + answer, _result = await assistant._run_agent_with_retries( + "Explain this field", None, asyncio.Queue() + ) + + assert answer == expected + assert len(calls) == 1 diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_automation_node_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_automation_node_tools.py index 0ba7a94baa..c2f0f2dcba 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_automation_node_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_automation_node_tools.py @@ -1,7 +1,19 @@ import pytest from baserow.contrib.automation.nodes.models import AutomationNode +from baserow.contrib.automation.nodes.operations import ( + UpdateAutomationNodeOperationType, +) from baserow.contrib.automation.nodes.service import AutomationNodeService +from baserow.contrib.automation.workflows.handler import AutomationWorkflowHandler +from baserow.core.cache import local_cache +from baserow.core.formula import resolve_formula +from baserow.core.formula.registries import formula_runtime_function_registry +from baserow.test_utils.pytest_conftest import FakeDispatchContext +from baserow_enterprise.assistant.tools.automation import agents as automation_agents +from baserow_enterprise.assistant.tools.automation.agents import ( + update_single_node_formulas, +) from baserow_enterprise.assistant.tools.automation.tools import ( add_nodes, create_workflows, @@ -15,6 +27,11 @@ TriggerNodeCreate, WorkflowCreate, ) +from baserow_enterprise.assistant.tools.automation.types.node import ( + AutomationFieldValue, +) +from baserow_enterprise.role.handler import RoleAssignmentHandler +from baserow_enterprise.role.models import Role from .utils import make_test_ctx @@ -326,27 +343,31 @@ def test_literal_values_with_apostrophes_stay_valid_formulas(data_fixture): user=user, workspace=workspace ) workflow = data_fixture.create_automation_workflow(automation=automation) - node = data_fixture.create_local_baserow_create_row_action_node( - user=user, workflow=workflow, service_kwargs={"table": table} - ) - - node_create = ActionNodeCreate( - ref="row1", - label="Create row", - previous_node_ref="trigger1", - type="create_row", - table_id=table.id, - row_id="Sales Managers' Week 3", - values=[{"field_id": field.id, "value": "Sales Managers' Week 3"}], + result = add_nodes( + make_test_ctx(user, workspace), + workflow_id=workflow.id, + nodes=[ + ActionNodeCreate( + ref="row1", + label="Create row", + previous_node_ref=str(workflow.get_trigger().id), + type="create_row", + table_id=table.id, + values=[{"field_id": field.id, "value": "Sales Managers' Week 3"}], + ) + ], + thought="Add an action with a fixed value", ) - node_create.apply_direct_values(node.service) - - service = node.service.specific - service.refresh_from_db() + service = AutomationNode.objects.get( + id=result["created_nodes"][0]["id"] + ).service.specific mapping = service.field_mappings.get(field_id=field.id) assert is_valid_formula(BaserowFormulaObject.to_formula(mapping.value)["formula"]) - assert is_valid_formula(BaserowFormulaObject.to_formula(service.row_id)["formula"]) + assert ( + resolve_formula(mapping.value, formula_runtime_function_registry, {}) + == "Sales Managers' Week 3" + ) # The workflow must remain duplicable, which is where the invalid formula # used to surface as a `BaserowFormulaSyntaxError`. @@ -412,3 +433,263 @@ def test_router_edge_conditions_stay_valid_formulas(data_fixture): ) assert duplicated.id != router_node.workflow.id + + +def test_node_create_types_fold_registered_names(): + trigger = TriggerNodeCreate( + ref="t", + label="T", + type="local_baserow_rows_created", + rows_triggers_settings={"table_id": 1}, + ) + assert trigger.type == "rows_created" + + action = ActionNodeCreate( + ref="a", + label="A", + previous_node_ref="t", + type="local_baserow_create_row", + table_id=1, + values=[AutomationFieldValue(field_id=1, value="x")], + ) + assert action.type == "create_row" + + +@pytest.mark.django_db(transaction=True) +def test_smtp_literal_recipients_are_stored(data_fixture): + user = data_fixture.create_user() + workflow = data_fixture.create_automation_workflow(user=user) + ctx = make_test_ctx(user, workflow.automation.workspace) + values = { + "to_emails": "recipient@example.com", + "cc_emails": "copy@example.com", + "bcc_emails": "audit@example.com", + "subject": "Review", + "body": "Ready", + } + result = add_nodes( + ctx, + workflow_id=workflow.id, + nodes=[ + ActionNodeCreate( + ref="email", + label="Send email", + type="smtp_email", + previous_node_ref=str(workflow.get_trigger().id), + **values, + ) + ], + thought="Create email action", + ) + node_id = result["created_nodes"][0]["id"] + service = AutomationNode.objects.get(id=node_id).service.specific + assert { + name: resolve_formula( + getattr(service, name), formula_runtime_function_registry, {} + ) + for name in values + } == values + + result = update_nodes( + ctx, + workflow_id=workflow.id, + nodes=[NodeUpdate(node_id=node_id, to_emails="updated@example.com")], + thought="Change the recipient only", + ) + + assert not result.get("errors") + service = AutomationNode.objects.get(id=node_id).service.specific + assert { + name: resolve_formula( + getattr(service, name), formula_runtime_function_registry, {} + ) + for name in values + } == {**values, "to_emails": "updated@example.com"} + + +@pytest.mark.django_db +@pytest.mark.parametrize("change_table", [True, False]) +def test_update_nodes_row_action_applies_table_and_literal_values( + data_fixture, row_action_node, change_table +): + user, node, original_field = row_action_node + workflow = node.workflow + database = node.service.specific.table.database + destination = ( + data_fixture.create_database_table(database=database) + if change_table + else original_field.table + ) + field = data_fixture.create_text_field(table=destination, name="Status") + row = destination.get_model().objects.create(id=42) + + result = update_nodes( + make_test_ctx(user, workflow.automation.workspace), + workflow_id=workflow.id, + nodes=[ + NodeUpdate( + node_id=node.id, + table_id=destination.id, + row_id=str(row.id), + values=[AutomationFieldValue(field_id=field.id, value="Reviewed")], + ) + ], + thought="Retarget the existing row action", + ) + + assert not result.get("errors") + assert result["updated_nodes"] == [{"node_id": node.id, "label": node.label}] + service = AutomationNode.objects.get(id=node.id).service.specific + assert service.table_id == destination.id + assert resolve_formula( + service.row_id, formula_runtime_function_registry, {} + ) == str(row.id) + mapping = service.field_mappings.get(field=field) + assert ( + resolve_formula(mapping.value, formula_runtime_function_registry, {}) + == "Reviewed" + ) + assert service.field_mappings.filter(field=original_field).exists() == ( + not change_table + ) + + service_type = service.get_type() + dispatch_context = FakeDispatchContext(context={}) + resolved_values = service_type.resolve_service_formulas(service, dispatch_context) + service_type.dispatch_data(service, resolved_values, dispatch_context) + row.refresh_from_db() + assert getattr(row, field.db_column) == "Reviewed" + + +@pytest.mark.django_db(transaction=True) +def test_single_node_formula_pass_accepts_a_create_payload(data_fixture, monkeypatch): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + automation, workflow_id = _create_test_workflow(data_fixture, user, workspace) + + ctx = make_test_ctx(user, workspace) + nodes = list_nodes(ctx, workflow_id=workflow_id, thought="inspect")["nodes"] + orm_node = AutomationNode.objects.get(id=int(nodes[1]["id"])).specific + + monkeypatch.setattr( + automation_agents, + "get_generate_formulas_tool", + lambda model_profile: ( + lambda formulas, context: {k: "concat('gen')" for k in formulas} + ), + ) + + node_create = ActionNodeCreate( + ref="email2", + label="Send Email", + previous_node_ref=str(nodes[0]["id"]), + type="smtp_email", + to_emails="$formula: the recipients from the trigger", + subject="Hello", + body="World", + ) + update_single_node_formulas(node_create, orm_node, ctx.deps.tool_helpers) + + service = AutomationNode.objects.get(id=orm_node.id).specific.service.specific + assert "gen" in str(service.to_emails) + + +@pytest.fixture +def row_action_node(data_fixture): + user = data_fixture.create_user() + workflow = data_fixture.create_automation_workflow(user=user) + database = data_fixture.create_database_application( + workspace=workflow.automation.workspace + ) + table = data_fixture.create_database_table(database=database) + field = data_fixture.create_text_field(table=table) + node = data_fixture.create_local_baserow_update_row_action_node( + workflow=workflow, + label="Update row", + service_kwargs={ + "table": table, + "row_id": "'1'", + "integration_args": { + "application": workflow.automation, + "authorized_user": user, + }, + }, + ) + node.service.specific.field_mappings.create(field=field, value="'Initial'") + return user, node, field + + +@pytest.mark.django_db +@pytest.mark.parametrize("property_name", ["values", "row_id"]) +def test_update_nodes_direct_values_require_update_permission( + enable_enterprise, enterprise_data_fixture, row_action_node, property_name +): + _, node, field = row_action_node + workflow = node.workflow + workspace = workflow.automation.workspace + user = enterprise_data_fixture.create_user() + enterprise_data_fixture.create_user_workspace( + user=user, workspace=workspace, permissions="NO_ACCESS" + ) + read_only_role = Role.objects.create(name="Read automation", workspace=workspace) + read_only_role.operations.set( + Role.objects.get(uid="BUILDER").operations.exclude( + name=UpdateAutomationNodeOperationType.type + ) + ) + RoleAssignmentHandler._init = False + RoleAssignmentHandler().assign_role( + user, workspace, role=read_only_role, scope=workflow.automation.application_ptr + ) + assert AutomationNodeService().get_node(user, node.id).id == node.id + payload = { + "values": [{"field_id": field.id, "value": "Updated"}], + "row_id": "2", + } + + result = update_nodes( + make_test_ctx(user, workspace), + workflow_id=workflow.id, + nodes=[NodeUpdate(node_id=node.id, **{property_name: payload[property_name]})], + thought="Update the action", + ) + + assert result["updated_nodes"] == [] + assert len(result["errors"]) == 1 + service = AutomationNode.objects.get(id=node.id).specific.service.specific + assert service.row_id["formula"] == "'1'" + assert service.field_mappings.get(field=field).value["formula"] == "'Initial'" + + +@pytest.mark.django_db +@pytest.mark.parametrize("property_name", ["values", "row_id"]) +def test_update_nodes_direct_values_refresh_test_run_clone( + row_action_node, property_name +): + user, node, field = row_action_node + workflow = node.workflow + handler = AutomationWorkflowHandler() + with local_cache.context(): + original_clone = handler._ensure_published_for_run(workflow) + payload = { + "values": [{"field_id": field.id, "value": "Updated"}], + "row_id": "2", + } + + result = update_nodes( + make_test_ctx(user, workflow.automation.workspace), + workflow_id=workflow.id, + nodes=[NodeUpdate(node_id=node.id, **{property_name: payload[property_name]})], + thought="Update the action", + ) + + assert not result.get("errors") + with local_cache.context(): + updated_clone = handler._ensure_published_for_run(workflow) + assert updated_clone.id != original_clone.id + cloned_node = updated_clone.automation_workflow_nodes.get(label=node.label) + service = cloned_node.service.specific + assert service.row_id["formula"] == ("'2'" if property_name == "row_id" else "'1'") + assert service.field_mappings.get(field=field).value["formula"] == ( + "'Updated'" if property_name == "values" else "'Initial'" + ) diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_builder_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_builder_tools.py index 5393b6cc64..10391e6e48 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_builder_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_builder_tools.py @@ -30,6 +30,7 @@ ButtonStyleOverride, CollectionElementCreate, DataSourceCreate, + DataSourceItem, DataSourceSort, DataSourceUpdate, DisplayElementCreate, @@ -345,7 +346,7 @@ def test_data_source_validation_errors(): @pytest.mark.parametrize("row_id", [None, "", " "]) def test_update_row_action_requires_a_row_id(row_id): - with pytest.raises(ValueError, match="'row_id' is required"): + with pytest.raises(ValueError, match="Missing: row_id"): ActionCreate( type="update_row", element="btn", @@ -533,6 +534,31 @@ def test_list_elements(data_fixture): assert result["elements"] == [] +@pytest.mark.django_db +def test_list_elements_on_populated_page(data_fixture): + # Regression: from_orm must only read attributes that survived b70bc968d. + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + builder = data_fixture.create_builder_application(user=user, workspace=workspace) + page = data_fixture.create_builder_page(builder=builder, name="Home", path="/home") + + data_fixture.create_builder_heading_element(page=page) + data_fixture.create_builder_form_container_element(page=page) + data_fixture.create_builder_table_element(page=page) + data_fixture.create_builder_column_element(page=page) + + ctx = make_test_ctx(user, workspace) + result = list_elements(ctx, page_id=page.id, thought="test") + + assert {el["type"] for el in result["elements"]} == { + "heading", + "form_container", + "table", + "column", + } + assert all(el["id"] for el in result["elements"]) + + @pytest.mark.django_db(transaction=True) def test_create_text_and_button(data_fixture): user = data_fixture.create_user() @@ -2633,3 +2659,18 @@ def test_setup_user_source_existing_table_creates_password_field(data_fixture): password_fields = [f for f in table_fields if isinstance(f, PasswordField)] assert len(password_fields) == 1 assert password_fields[0].name == "Password" + + +def test_data_source_type_folds_registered_name(): + ds = DataSourceCreate( + ref="d", name="Products", type="local_baserow_list_rows", table_id=1 + ) + assert ds.type == "list_rows" + + +def test_data_source_dedup_matches_the_registered_type(): + ds = DataSourceCreate(ref="d", name="Products", type="list_rows", table_id=5) + existing = DataSourceItem( + id=1, name="Existing", type="local_baserow_list_rows", table_id=5 + ) + assert ds.matches_existing(existing) diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_core_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_core_tools.py index 49c3a9cc84..d8de321a7e 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_core_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_core_tools.py @@ -4,8 +4,12 @@ from baserow_enterprise.assistant.tools.core.tools import ( create_builders, list_builders, + update_builder, +) +from baserow_enterprise.assistant.tools.core.types import ( + BuilderItemCreate, + BuilderUpdate, ) -from baserow_enterprise.assistant.tools.core.types import BuilderItemCreate from .utils import make_test_ctx @@ -158,3 +162,66 @@ def test_create_database_ignores_theme(data_fixture): assert len(result["created_builders"]) == 1 assert result["created_builders"][0]["type"] == "database" + + +@pytest.mark.django_db +def test_update_builder_renames_an_application(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application( + workspace=workspace, name="Old Name" + ) + + ctx = make_test_ctx(user, workspace) + result = update_builder( + ctx, + builder_id=database.id, + update=BuilderUpdate(name="New Name"), + thought="rename", + ) + + assert result == {"id": database.id, "name": "New Name"} + database.refresh_from_db() + assert database.name == "New Name" + + +@pytest.mark.django_db +def test_update_builder_sets_the_login_page(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + builder = data_fixture.create_builder_application( + user=user, workspace=workspace, name="Portal" + ) + page = data_fixture.create_builder_page( + builder=builder, name="Login", path="/login" + ) + + ctx = make_test_ctx(user, workspace) + result = update_builder( + ctx, + builder_id=builder.id, + update=BuilderUpdate(login_page_id=page.id), + thought="set login page", + ) + + assert result["login_page_id"] == page.id + builder.refresh_from_db() + assert builder.login_page_id == page.id + + +@pytest.mark.django_db +def test_update_builder_with_nothing_to_change_is_a_no_op(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application( + workspace=workspace, name="Unchanged" + ) + + ctx = make_test_ctx(user, workspace) + result = update_builder( + ctx, builder_id=database.id, update=BuilderUpdate(), thought="no-op" + ) + + assert result["name"] == "Unchanged" + database.refresh_from_db() + assert database.name == "Unchanged" diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_rows_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_rows_tools.py index 306c06289a..65c19cb975 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_rows_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_rows_tools.py @@ -1,14 +1,223 @@ import pytest +from pydantic import ValidationError +from pydantic_ai import ModelRetry from baserow.contrib.database.rows.handler import RowHandler +from baserow_enterprise.assistant.agents import dynamic_toolset +from baserow_enterprise.assistant.tools.database import tools as database_tools from baserow_enterprise.assistant.tools.database.tools import ( + create_fields, + delete_fields, list_rows, load_row_tools, + update_fields, +) +from baserow_enterprise.assistant.tools.database.types import ( + FieldItemCreate, + FieldItemUpdate, ) from .utils import make_test_ctx +@pytest.mark.django_db +def test_create_rows_rejects_empty_payload_and_accepts_corrected_call(data_fixture): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + name_field = data_fixture.create_text_field(table=table, name="Name", primary=True) + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["create"], thought="Prepare row creation") + tool = ctx.deps.dynamic_tools[0] + model = table.get_model() + + with pytest.raises(ModelRetry, match="Nothing was changed"): + tool.function(rows=[], thought="Create rows") + + assert model.objects.count() == 0 + + arguments = tool.function_schema.validator.validate_python( + {"rows": [{"Name": "Created after retry"}], "thought": "Create the row"} + ) + result = tool.function(**arguments) + + assert list(model.objects.values_list("id", name_field.db_column)) == [ + (result["created_row_ids"][0], "Created after retry") + ] + + +@pytest.mark.django_db +def test_reload_row_tools_uses_new_schema_without_duplicate_tools(data_fixture): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + data_fixture.create_text_field(table=table, name="Name", primary=True) + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["create", "delete"], thought="Prepare rows") + old_create, delete = ctx.deps.dynamic_tools + arguments = { + "rows": [{"Name": "Batch 46", "Status": "Reviewed"}], + "thought": "Create a row with its status", + } + + with pytest.raises(ValidationError, match="Status"): + old_create.function_schema.validator.validate_python(arguments) + + create_fields( + ctx, + table_id=table.id, + fields=[FieldItemCreate(name="Status", type="text")], + thought="Add a status field", + ) + load_row_tools(ctx, [table.id], ["create"], thought="Refresh row creation") + toolset = dynamic_toolset(ctx) + create = toolset.tools[f"create_rows_in_table_{table.id}"] + + assert len(toolset.tools) == 2 + assert toolset.tools[delete.name] is delete + assert create is not old_create + result = create.function( + **create.function_schema.validator.validate_python(arguments) + ) + + status_field = table.field_set.get(name="Status") + assert list( + table.get_model().objects.values_list("id", status_field.db_column) + ) == [(result["created_row_ids"][0], "Reviewed")] + + +@pytest.mark.django_db +def test_create_fields_refreshes_loaded_row_tools(data_fixture): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + data_fixture.create_text_field(table=table, name="Name", primary=True) + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["create", "delete"], thought="Prepare rows") + stale_create, delete = ctx.deps.dynamic_tools + + create_fields( + ctx, + table_id=table.id, + fields=[FieldItemCreate(name="Status", type="text")], + thought="Add a status field", + ) + + create, unchanged_delete = ctx.deps.dynamic_tools + assert create is not stale_create + assert unchanged_delete is delete + result = create.function( + **create.function_schema.validator.validate_python( + { + "rows": [{"Name": "Batch 46", "Status": "Reviewed"}], + "thought": "Create a row with its status", + } + ) + ) + + status_field = table.field_set.get(name="Status") + assert list( + table.get_model().objects.values_list("id", status_field.db_column) + ) == [(result["created_row_ids"][0], "Reviewed")] + + +@pytest.mark.django_db +def test_update_fields_refreshes_loaded_row_tools(data_fixture): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + field = data_fixture.create_text_field(table=table, name="Name", primary=True) + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["update"], thought="Prepare updates") + + update_fields( + ctx, + fields=[FieldItemUpdate(field_id=field.id, name="Title")], + thought="Rename the field", + ) + + update = ctx.deps.dynamic_tools[0] + with pytest.raises(ValidationError, match="Name"): + update.function_schema.validator.validate_python( + {"rows": [{"id": 1, "Name": "Renamed"}], "thought": "Update the row"} + ) + assert update.function_schema.validator.validate_python( + {"rows": [{"id": 1, "Title": "Renamed"}], "thought": "Update the row"} + ) + + +@pytest.mark.django_db +def test_delete_fields_refreshes_loaded_row_tools(data_fixture): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + data_fixture.create_text_field(table=table, name="Name", primary=True) + notes = data_fixture.create_long_text_field(table=table, name="Notes") + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["create"], thought="Prepare rows") + + delete_fields(ctx, field_ids=[notes.id], thought="Drop the notes field") + + create = ctx.deps.dynamic_tools[0] + with pytest.raises(ValidationError, match="Notes"): + create.function_schema.validator.validate_python( + { + "rows": [{"Name": "Batch 46", "Notes": "Gone"}], + "thought": "Create a row", + } + ) + + +@pytest.mark.django_db +def test_create_fields_refreshes_row_tools_of_the_linked_table(data_fixture): + user = data_fixture.create_user() + database = data_fixture.create_database_application(user=user) + table_a = data_fixture.create_database_table(user=user, database=database) + table_b = data_fixture.create_database_table(user=user, database=database) + data_fixture.create_text_field(table=table_a, name="Name", primary=True) + data_fixture.create_text_field(table=table_b, name="Title", primary=True) + ctx = make_test_ctx(user, database.workspace) + load_row_tools(ctx, [table_b.id], ["create"], thought="Prepare rows in B") + + create_fields( + ctx, + table_id=table_a.id, + fields=[ + FieldItemCreate(name="B rows", type="link_row", linked_table=table_b.id) + ], + thought="Link A to B", + ) + + reverse_field = table_b.field_set.exclude(primary=True).get() + create_b = ctx.deps.dynamic_tools[0] + assert create_b.function_schema.validator.validate_python( + { + "rows": [{"Title": "Row B", reverse_field.name: []}], + "thought": "Create a row in B", + } + ) + + +@pytest.mark.django_db +def test_field_change_survives_a_failing_row_tool_refresh(data_fixture, monkeypatch): + user = data_fixture.create_user() + table = data_fixture.create_database_table(user=user) + data_fixture.create_text_field(table=table, name="Name", primary=True) + ctx = make_test_ctx(user, table.database.workspace) + load_row_tools(ctx, [table.id], ["create"], thought="Prepare rows") + stale_create = ctx.deps.dynamic_tools[0] + + def _raise(*args, **kwargs): + raise RuntimeError("cannot build row tools") + + monkeypatch.setattr(database_tools, "_build_row_tools", _raise) + result = create_fields( + ctx, + table_id=table.id, + fields=[FieldItemCreate(name="Status", type="text")], + thought="Add a status field", + ) + + assert result["created_fields"][0]["name"] == "Status" + assert table.field_set.filter(name="Status").exists() + assert ctx.deps.dynamic_tools[0] is stale_create + + def _create_simple_database_with_linked_tables_and_rows(data_fixture): user = data_fixture.create_user() table_a, table_b, link_a_to_b = data_fixture.create_two_linked_tables(user=user) diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_table_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_table_tools.py index fd47e1b038..c36521e898 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_table_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_table_tools.py @@ -1,6 +1,7 @@ from unittest.mock import MagicMock, patch import pytest +from pydantic_ai import ModelRetry from baserow.contrib.database.fields.models import FormulaField from baserow.contrib.database.formula.registries import formula_function_registry @@ -310,7 +311,7 @@ def test_generate_formula_no_save(data_fixture): model_profile.get_settings.return_value = model_settings with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model" + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model" ) as mock_agent: mock_agent.return_value = mock_result @@ -352,7 +353,7 @@ def test_generate_formula_create_new_field(data_fixture): ) with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model" + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model" ) as mock_agent: mock_agent.return_value = mock_result @@ -404,7 +405,7 @@ def test_generate_formula_update_existing_formula_field(data_fixture): ) with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model" + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model" ) as mock_agent: mock_agent.return_value = mock_result @@ -456,7 +457,7 @@ def test_generate_formula_replace_non_formula_field(data_fixture): ) with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model" + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model" ) as mock_agent: mock_agent.return_value = mock_result @@ -510,14 +511,14 @@ def test_generate_formula_invalid_formula(data_fixture): ) with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model" + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model" ) as mock_agent: mock_agent.return_value = mock_result ctx = make_test_ctx(user, workspace) - # Verify exception is raised - with pytest.raises(Exception) as exc_info: + # A generation failure is recoverable: it must become a retry prompt. + with pytest.raises(ModelRetry) as exc_info: generate_formula( ctx, thought="test", @@ -526,7 +527,6 @@ def test_generate_formula_invalid_formula(data_fixture): save_to_field=True, ) - assert "Error generating formula:" in str(exc_info.value) assert "Formula syntax error: invalid expression" in str(exc_info.value) # Verify no field was created @@ -558,7 +558,7 @@ def mock_run_sync(_agent, prompt, **kwargs): return mock_result with patch( - "baserow_enterprise.assistant.tools.database.tools.run_agent_sync_with_model", + "baserow_enterprise.assistant.tools.database.agents.run_agent_sync_with_model", side_effect=mock_run_sync, ): ctx = make_test_ctx(user, workspace) diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_views_tools.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_views_tools.py index 4ac15bd2e7..fdf88abe27 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_views_tools.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_database_views_tools.py @@ -1,4 +1,5 @@ import pytest +from pydantic_ai import ModelRetry from baserow.contrib.database.views.models import View, ViewFilter from baserow_enterprise.assistant.tools.database.tools import ( @@ -72,7 +73,8 @@ def test_create_grid_view(data_fixture): assert len(response["created_views"]) == 1 assert response["created_views"][0]["name"] == "Grid View" - assert View.objects.filter(name="Grid View").exists() + grid = View.objects.get(name="Grid View").specific + assert grid.row_height_size == "medium" @pytest.mark.django_db @@ -849,3 +851,118 @@ def test_create_multiple_select_is_none_of_filter(data_fixture): assert ViewFilter.objects.filter( view=view, field=field, type="multiple_select_has_not" ).exists() + + +@pytest.mark.django_db +def test_create_views_bad_cover_field_asks_the_model_to_retry(data_fixture): + # A bad field id must come back as a retry prompt naming the valid fields. + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application(workspace=workspace) + table = data_fixture.create_database_table(database=database, name="Photos") + data_fixture.create_text_field(table=table, name="Caption") + + ctx = make_test_ctx(user, workspace) + + with pytest.raises(ModelRetry) as exc_info: + create_views( + ctx, + table_id=table.id, + views=[ViewItemCreate(name="Gallery", type="gallery", cover_field_id=0)], + thought="test", + ) + + message = str(exc_info.value) + assert "Gallery" in message + assert "Caption" in message + + +@pytest.mark.django_db +def test_create_views_wrong_field_type_asks_the_model_to_retry(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application(workspace=workspace) + table = data_fixture.create_database_table(database=database, name="Tasks") + text_field = data_fixture.create_text_field(table=table, name="Status") + + ctx = make_test_ctx(user, workspace) + + with pytest.raises(ModelRetry) as exc_info: + create_views( + ctx, + table_id=table.id, + views=[ + ViewItemCreate( + name="Board", type="kanban", column_field_id=text_field.id + ) + ], + thought="test", + ) + + assert "Single Select" in str(exc_info.value) + + +@pytest.mark.django_db +@pytest.mark.parametrize("include_compatible", [True, False]) +def test_create_form_view_skips_incompatible_fields(data_fixture, include_compatible): + """Report only saved inputs, including when every requested field is skipped.""" + + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application(workspace=workspace) + table = data_fixture.create_database_table(database=database) + name = data_fixture.create_text_field(table=table, name="Name", primary=True) + formula = data_fixture.create_formula_field( + table=table, name="Computed", formula="'x'" + ) + supported_options = ( + [ + { + "field_id": name.id, + "name": "Your name", + "description": "Enter your name", + "required": True, + "order": 1, + } + ] + if include_compatible + else [] + ) + + ctx = make_test_ctx(user, workspace) + response = create_views( + ctx, + thought="test", + table_id=table.id, + views=[ + ViewItemCreate( + name="Signup", + public=False, + type="form", + field_options=[ + *[FormFieldOption(**options) for options in supported_options], + FormFieldOption( + field_id=formula.id, + name="Computed", + description="", + required=False, + order=2, + ), + ], + ) + ], + ) + + created = response["created_views"][0] + assert "Computed" in created["skipped_fields"] + assert "Not shown on the form" in created["skipped_fields"] + assert created["field_options"] == supported_options + form = View.objects.get(name="Signup").specific + assert ( + list( + form.active_field_options.values( + "field_id", "name", "description", "required", "order" + ) + ) + == supported_options + ) diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_formula_agent.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_formula_agent.py new file mode 100644 index 0000000000..c64d81a1f1 --- /dev/null +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_assistant_formula_agent.py @@ -0,0 +1,367 @@ +""" +Unit tests for the formula validation tool used by the formula sub-agent. + +The contract that matters is the exception *class*: pydantic-ai turns only +``ModelRetry`` into a retry prompt, so any other exception escaping this tool +aborts the user's entire turn instead of letting the agent fix its formula. +""" + +from types import SimpleNamespace + +import pytest +from pydantic_ai import ModelRetry +from pydantic_ai.messages import ( + ModelResponse, + RetryPromptPart, + ToolCallPart, + ToolReturnPart, +) +from pydantic_ai.models.function import FunctionModel + +from baserow_enterprise.assistant import model_profiles +from baserow_enterprise.assistant.tools.database import agents as database_agents +from baserow_enterprise.assistant.tools.database.agents import ( + GET_FORMULA_TYPE_TOOL_NAME, + FormulaGenerationResult, + _type_mismatch_hint, + _verdict_must_be_backed_by_validation, + get_formula_type_tool, +) + +from .utils import create_fake_tool_helpers + + +@pytest.fixture +def formula_env(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application(workspace=workspace) + table = data_fixture.create_database_table(database=database, name="Orders") + data_fixture.create_text_field(table=table, name="Customer") + data_fixture.create_number_field(table=table, name="Amount") + return get_formula_type_tool(user, workspace), table + + +@pytest.mark.django_db +def test_valid_formula_returns_its_type(formula_env): + validate, table = formula_env + + assert validate(table.id, "Label", "concat(field('Customer'), '!')") == "text" + + +@pytest.mark.django_db +@pytest.mark.parametrize( + "formula,expected_in_message", + [ + ("or(true, true, true)", "or"), + ("{Amount} + 1", "{"), + ("weekday(field('Amount'))", "weekday"), + ("field('Nonexistent')", "Nonexistent"), + ], + ids=["or-arity", "curly-braces", "unknown-function", "unknown-field"], +) +def test_invalid_formula_asks_the_model_to_retry( + formula_env, formula, expected_in_message +): + validate, table = formula_env + + with pytest.raises(ModelRetry) as exc_info: + validate(table.id, "Broken", formula) + + assert expected_in_message in str(exc_info.value) + + +@pytest.mark.django_db +def test_rejection_lists_the_available_fields(formula_env): + validate, table = formula_env + + with pytest.raises(ModelRetry) as exc_info: + validate(table.id, "Broken", "field('Nonexistent')") + + message = str(exc_info.value) + assert "Customer" in message + assert "Amount" in message + + +@pytest.mark.django_db +def test_unknown_table_asks_the_model_to_retry_with_valid_ids(formula_env): + validate, table = formula_env + + with pytest.raises(ModelRetry) as exc_info: + validate(table.id + 999, "Broken", "field('Customer')") + + assert str(table.id) in str(exc_info.value) + + +@pytest.mark.django_db +def test_table_outside_the_workspace_is_not_leaked(data_fixture): + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + other_table = data_fixture.create_database_table(name="Secret") + + validate = get_formula_type_tool(user, workspace) + + with pytest.raises(ModelRetry): + validate(other_table.id, "Broken", "field('Customer')") + + +@pytest.mark.django_db +@pytest.mark.parametrize( + "formula,conversion", + [ + ("day(field('Customer'))", "todate"), + ("upper(field('Amount'))", "totext"), + ], + ids=["text-into-date-slot", "number-into-text-slot"], +) +def test_type_mismatch_rejection_carries_a_conversion_hint( + formula_env, formula, conversion +): + """Pins the _USABLE_TYPES regex to the compiler wording in ast/tree.py.""" + + validate, table = formula_env + + with pytest.raises(ModelRetry) as exc_info: + validate(table.id, "Broken", formula) + + message = str(exc_info.value) + assert "argument type mismatch" in message + assert conversion in message + + +def test_no_usable_type_hint_says_restructure(): + hint = _type_mismatch_hint( + "argument number 1 given to function x was of type text " + "but there are no possible types usable here" + ) + assert "restructure" in hint + + +def _tool_call(call_id: str, formula: str) -> ToolCallPart: + return ToolCallPart( + tool_name=GET_FORMULA_TYPE_TOOL_NAME, + args={"table_id": 1, "field_name": "F", "formula": formula}, + tool_call_id=call_id, + ) + + +def _rejection(call_id: str) -> RetryPromptPart: + return RetryPromptPart( + content="Invalid formula", + tool_name=GET_FORMULA_TYPE_TOOL_NAME, + tool_call_id=call_id, + ) + + +def _acceptance(call_id: str) -> ToolReturnPart: + return ToolReturnPart( + tool_name=GET_FORMULA_TYPE_TOOL_NAME, content="number", tool_call_id=call_id + ) + + +def _run_ctx(*parts) -> SimpleNamespace: + return SimpleNamespace(messages=[SimpleNamespace(parts=list(parts))]) + + +def _verdict(formula: str = "field('Amount') * 2", valid: bool = True): + return FormulaGenerationResult( + table_id=1, + field_name="F", + formula=formula, + formula_type="number", + is_formula_valid=valid, + error_message="" if valid else "cannot be expressed", + ) + + +def test_valid_verdict_without_any_tool_call_is_sent_back(): + with pytest.raises(ModelRetry, match="never accepted"): + _verdict_must_be_backed_by_validation(_run_ctx(), _verdict()) + + +def test_valid_verdict_naming_a_rejected_formula_is_sent_back(): + ctx = _run_ctx(_tool_call("c1", "field('Amount') * 2"), _rejection("c1")) + + with pytest.raises(ModelRetry, match="never accepted"): + _verdict_must_be_backed_by_validation(ctx, _verdict()) + + +def test_valid_verdict_matches_the_accepted_formula_and_field(): + ctx = _run_ctx(_tool_call("c1", "field('Amount') * 2"), _acceptance("c1")) + output = _verdict() + + assert _verdict_must_be_backed_by_validation(ctx, output) is output + + +@pytest.mark.parametrize("changed_field", ["table_id", "field_name"]) +def test_valid_verdict_must_match_the_validated_field(changed_field): + ctx = _run_ctx(_tool_call("c1", "field('Amount') * 2"), _acceptance("c1")) + output = _verdict() + if changed_field == "table_id": + output.table_id = 2 + else: + output.field_name = "Amount" + + with pytest.raises(ModelRetry, match="never accepted"): + _verdict_must_be_backed_by_validation(ctx, output) + + +@pytest.mark.parametrize( + "validated_formula,returned_formula", + [ + ("field('Amount') *\n2", "field('Amount') * 2"), + ("field('Two Spaces')", "field('Two Spaces')"), + ("'two spaces'", "'two spaces'"), + ('"two spaces"', '"two spaces"'), + ], + ids=["formatting", "field-reference", "single-quoted", "double-quoted"], +) +def test_valid_verdict_requires_the_exact_validated_formula( + validated_formula, returned_formula +): + ctx = _run_ctx(_tool_call("c1", validated_formula), _acceptance("c1")) + + with pytest.raises(ModelRetry, match="never accepted"): + _verdict_must_be_backed_by_validation(ctx, _verdict(returned_formula)) + + +def test_impossible_verdict_after_one_rejection_is_sent_back(): + ctx = _run_ctx(_tool_call("c1", "day(field('Customer'))"), _rejection("c1")) + + with pytest.raises(ModelRetry, match="materially different"): + _verdict_must_be_backed_by_validation(ctx, _verdict(valid=False)) + + +def test_impossible_verdict_after_retrying_the_same_formula_is_sent_back(): + ctx = _run_ctx( + _tool_call("c1", "day( field('Customer') )"), + _rejection("c1"), + _tool_call("c2", "day(\nfield('Customer')\n)"), + _rejection("c2"), + ) + + with pytest.raises(ModelRetry, match="materially different"): + _verdict_must_be_backed_by_validation(ctx, _verdict(valid=False)) + + +def test_impossible_verdict_backed_by_two_different_rejections_passes(): + ctx = _run_ctx( + _tool_call("c1", "day(field('Customer'))"), + _rejection("c1"), + _tool_call("c2", "todate(field('Customer'), 'YYYY')"), + _rejection("c2"), + ) + output = _verdict(valid=False) + + assert _verdict_must_be_backed_by_validation(ctx, output) is output + + +@pytest.mark.parametrize("raw_table_id", [1, "1.0"]) +def test_formula_generation_reuses_the_request_profile_and_owns_its_model( + monkeypatch, raw_table_id +): + lifecycle = [] + requested_models = [] + observed_settings = [] + + def get_formula_type(table_id: int, field_name: str, formula: str) -> str: + assert (table_id, field_name, formula) == (1, "F", "field('Amount') * 2") + return "number" + + def respond(messages, info): + observed_settings.append(info.model_settings) + if any( + isinstance(part, ToolReturnPart) + for message in messages + for part in message.parts + ): + return ModelResponse( + parts=[ + ToolCallPart( + tool_name="final_result", + args=_verdict().model_dump(), + tool_call_id="result", + ) + ] + ) + return ModelResponse( + parts=[ + ToolCallPart( + tool_name=GET_FORMULA_TYPE_TOOL_NAME, + args={ + "table_id": raw_table_id, + "field_name": "F", + "formula": "field('Amount') * 2", + }, + tool_call_id="validation", + ) + ] + ) + + class LifecycleModel(FunctionModel): + async def __aenter__(self): + lifecycle.append("entered") + return await super().__aenter__() + + async def __aexit__(self, *args): + lifecycle.append("exited") + return await super().__aexit__(*args) + + model = LifecycleModel(respond) + + def create_model(model_string): + requested_models.append(model_string) + return model + + monkeypatch.setattr(model_profiles, "RetryingModel", create_model) + monkeypatch.setattr( + database_agents, + "get_formula_type_tool", + lambda user, workspace: get_formula_type, + ) + profile = model_profiles.ResolvedAssistantModelProfile( + model_string="openai:gpt-4.1-mini", + source="explicit", + workspace=None, + database_model=None, + ) + + result = database_agents.run_formula_generation(None, None, "Generate", profile) + + assert result.output == _verdict() + assert requested_models == [profile.model_string] + assert observed_settings == [profile.get_settings(model_profiles.UTILITY)] * 2 + assert lifecycle == ["entered", "exited"] + validation_call = next( + part + for message in result.all_messages() + for part in message.parts + if isinstance(part, ToolCallPart) + and part.tool_name == GET_FORMULA_TYPE_TOOL_NAME + ) + assert validation_call.args_as_dict()["table_id"] == raw_table_id + + +@pytest.mark.django_db +def test_formula_fixer_contains_generator_failures(data_fixture, monkeypatch): + """The fixer runs inside another except handler, so it must never raise.""" + + from pydantic_ai.exceptions import UnexpectedModelBehavior + + user = data_fixture.create_user() + workspace = data_fixture.create_workspace(user=user) + database = data_fixture.create_database_application(workspace=workspace) + table = data_fixture.create_database_table(database=database, name="Orders") + data_fixture.create_text_field(table=table, name="Customer", primary=True) + + def raise_retries_exhausted(*args, **kwargs): + raise UnexpectedModelBehavior("Exceeded maximum output retries (3)") + + monkeypatch.setattr( + database_agents, "run_agent_sync_with_model", raise_retries_exhausted + ) + + tool_helpers = create_fake_tool_helpers() + fix_formula = database_agents.make_formula_fixer(user, workspace, tool_helpers) + + assert fix_formula(table, "Total", "field('Missing') *") is None diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_pydantic_ai_contract.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_pydantic_ai_contract.py index 5c3f989706..0468477420 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_pydantic_ai_contract.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_pydantic_ai_contract.py @@ -281,8 +281,22 @@ def track_tool(x: str) -> str: tool_calls.append(x) return "tracked" + def get_formula_type(table_id: int, field_name: str, formula: str) -> str: + return "text" + def func(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: + # The output validator only accepts formulas get_formula_type validated. if len(messages) == 1: + return ModelResponse( + parts=[ + ToolCallPart( + tool_name="get_formula_type", + args={"table_id": 1, "field_name": "f", "formula": "'ok'"}, + tool_call_id="0", + ), + ] + ) + if len(messages) == 3: return ModelResponse( parts=[ ToolCallPart( @@ -303,7 +317,7 @@ def func(messages: list[ModelMessage], info: AgentInfo) -> ModelResponse: ) return ModelResponse(parts=[TextPart(content="done")]) - toolset = FunctionToolset([Tool(track_tool)]) + toolset = FunctionToolset([Tool(track_tool), Tool(get_formula_type)]) result = formula_generation_agent.run_sync( "generate a formula", model=FunctionModel(func), diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_retrying_model.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_retrying_model.py index 3d6352f89b..95ef50b4c2 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_retrying_model.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_retrying_model.py @@ -31,6 +31,17 @@ def test_generic_error_not_retryable(self): exc = ValueError("something went wrong") assert _is_transient_provider_error(exc) is False + def test_rate_limit_http_error_is_retryable(self): + from pydantic_ai.exceptions import ModelHTTPError + + exc = ModelHTTPError( + status_code=429, + model_name="test-model", + body={"error": {"message": "Rate limit exceeded"}}, + ) + + assert _is_transient_provider_error(exc) is True + def _make_retrying(inner_mock, **kwargs): """Create a RetryingModel with a pre-resolved mock as the wrapped model.""" @@ -78,6 +89,37 @@ async def test_request_retries_on_transient_error(): _assert_balanced_model_scope(inner) +@pytest.mark.asyncio +async def test_request_uses_provider_retry_after_for_rate_limit(): + from unittest.mock import AsyncMock, MagicMock, patch + + from pydantic_ai.exceptions import ModelHTTPError + from pydantic_ai.messages import ModelResponse, TextPart + + rate_limit_error = ModelHTTPError( + status_code=429, + model_name="test-model", + body={"error": {"message": "Rate limit exceeded"}}, + headers={"Retry-After": "2"}, + ) + response = ModelResponse(parts=[TextPart(content="hello")]) + inner = MagicMock() + inner.request = AsyncMock(side_effect=[rate_limit_error, response]) + model = _make_retrying(inner, base_delay=0.01, max_delay=1.0) + + with patch( + "baserow_enterprise.assistant.retrying_model.asyncio.sleep", + new_callable=AsyncMock, + ) as mock_sleep: + result = await model.request( + [], None, ModelRequestParameters(function_tools=[], output_tools=[]) + ) + + assert result == response + assert inner.request.call_count == 2 + mock_sleep.assert_awaited_once_with(1.0) + + @pytest.mark.asyncio async def test_request_raises_non_transient_error(): """RetryingModel.request should not retry non-transient errors.""" diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_telemetry.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_telemetry.py index 909dfdea11..55d91a91f4 100644 --- a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_telemetry.py +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_telemetry.py @@ -1,3 +1,4 @@ +import asyncio import json import sys from unittest.mock import MagicMock, patch @@ -9,6 +10,7 @@ from baserow_enterprise.assistant import telemetry from baserow_enterprise.assistant.models import AssistantChat from baserow_enterprise.assistant.telemetry import ( + AssistantTraceOutcome, PosthogSpanProcessor, PosthogTracingCallback, _pydantic_messages_to_posthog, @@ -40,10 +42,11 @@ def test_trace_context_manager_success( callback = PosthogTracingCallback() - with callback.trace(assistant_chat_fixture, "Hello"): + with callback.trace(assistant_chat_fixture, "Hello") as tracer: assert callback.trace_id is not None assert callback.span_id is not None assert callback.user_id == str(assistant_chat_fixture.user_id) + tracer.set_trace_output("Hi there") # Verify trace event captured mock_posthog.capture.assert_called_once() @@ -64,7 +67,8 @@ def test_trace_context_manager_success( assert props["$ai_latency"] >= 0 assert props["$ai_is_error"] is False assert props["$ai_input_state"] == {"user_message": "Hello"} - assert props["$ai_output_state"] is None + assert props["$ai_output_state"] == {"answer": "Hi there"} + assert props["assistant_outcome"] == AssistantTraceOutcome.ANSWERED @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") def test_trace_context_manager_exception( @@ -85,7 +89,10 @@ def test_trace_context_manager_exception( call_args = mock_posthog.capture.call_args assert call_args is not None assert call_args.kwargs["event"] == "$ai_trace" - assert call_args.kwargs["properties"]["$ai_is_error"] is True + props = call_args.kwargs["properties"] + assert props["$ai_is_error"] is True + assert props["$ai_output_state"] == "Test error" + assert props["assistant_outcome"] == AssistantTraceOutcome.ERROR @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") def test_trace_with_output(self, mock_get_client, assistant_chat_fixture): @@ -103,6 +110,118 @@ def test_trace_with_output(self, mock_get_client, assistant_chat_fixture): props = call_args.kwargs["properties"] assert props["$ai_output_state"] == {"answer": "The answer is 42"} + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") + def test_trace_merges_tool_calls_into_the_answer( + self, mock_get_client, assistant_chat_fixture + ): + """Tool names recorded during the run are merged into the output.""" + + mock_posthog = MagicMock() + mock_get_client.return_value = mock_posthog + + callback = PosthogTracingCallback() + + with callback.trace(assistant_chat_fixture, "Hello") as tracer: + _tool_calls.get().append("list_tables") + _tool_calls.get().append("create_rows") + tracer.set_trace_output("Done") + + props = mock_posthog.capture.call_args.kwargs["properties"] + assert props["$ai_output_state"] == { + "answer": "Done", + "tool_calls": ["list_tables", "create_rows"], + } + + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") + def test_trace_exception_keeps_the_message_without_tool_calls( + self, mock_get_client, assistant_chat_fixture + ): + """The error path stays a plain string, so no tool names are merged.""" + + mock_posthog = MagicMock() + mock_get_client.return_value = mock_posthog + + callback = PosthogTracingCallback() + + with pytest.raises(ValueError): + with callback.trace(assistant_chat_fixture, "Hello"): + _tool_calls.get().append("list_tables") + raise ValueError("Boom") + + props = mock_posthog.capture.call_args.kwargs["properties"] + assert props["$ai_output_state"] == "Boom" + assert props["assistant_outcome"] == AssistantTraceOutcome.ERROR + + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") + def test_trace_without_answer_is_an_error( + self, mock_get_client, assistant_chat_fixture + ): + """A run that ends normally without an answer is a silent failure.""" + + mock_posthog = MagicMock() + mock_get_client.return_value = mock_posthog + + callback = PosthogTracingCallback() + + with callback.trace(assistant_chat_fixture, "Hello"): + _tool_calls.get().append("create_rows") + + props = mock_posthog.capture.call_args.kwargs["properties"] + assert props["$ai_is_error"] is True + assert props["assistant_outcome"] == AssistantTraceOutcome.NO_ANSWER + assert props["$ai_output_state"] == { + "status": "no_answer", + "tool_calls": ["create_rows"], + } + + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") + def test_trace_cancelled_by_user_is_not_an_error( + self, mock_get_client, assistant_chat_fixture + ): + """Pressing stop is a legitimate no-answer run, not a failure.""" + + mock_posthog = MagicMock() + mock_get_client.return_value = mock_posthog + + callback = PosthogTracingCallback() + + with pytest.raises(asyncio.CancelledError): + with callback.trace( + assistant_chat_fixture, "Hello", cancelled_by_user=lambda: True + ): + _tool_calls.get().append("list_tables") + raise asyncio.CancelledError() + + props = mock_posthog.capture.call_args.kwargs["properties"] + assert props["$ai_is_error"] is False + assert props["assistant_outcome"] == AssistantTraceOutcome.CANCELLED + assert props["$ai_output_state"] == { + "status": "cancelled", + "tool_calls": ["list_tables"], + } + + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") + def test_trace_interrupted_without_user_cancel( + self, mock_get_client, assistant_chat_fixture + ): + """A dropped run is reported apart from a deliberate cancel.""" + + mock_posthog = MagicMock() + mock_get_client.return_value = mock_posthog + + callback = PosthogTracingCallback() + + with pytest.raises(asyncio.CancelledError): + with callback.trace( + assistant_chat_fixture, "Hello", cancelled_by_user=lambda: False + ): + raise asyncio.CancelledError() + + props = mock_posthog.capture.call_args.kwargs["properties"] + assert props["$ai_is_error"] is False + assert props["assistant_outcome"] == AssistantTraceOutcome.INTERRUPTED + assert props["$ai_output_state"] == {"status": "interrupted"} + @patch("baserow_enterprise.assistant.telemetry.get_posthog_client") def test_trace_sets_and_clears_context_var( self, mock_get_client, assistant_chat_fixture diff --git a/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_toolset.py b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_toolset.py new file mode 100644 index 0000000000..c032eee9fd --- /dev/null +++ b/enterprise/backend/tests/baserow_enterprise_tests/assistant/test_toolset.py @@ -0,0 +1,219 @@ +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest +from pydantic import ValidationError +from pydantic_ai import ModelRetry, RunContext +from pydantic_ai.models.test import TestModel +from pydantic_ai.toolsets import FunctionToolset +from pydantic_ai.usage import RunUsage + +from baserow_enterprise.assistant.tools.automation.types.node import ( + ActionNodeCreate, + TriggerNodeCreate, +) +from baserow_enterprise.assistant.tools.builder.types.data_source import ( + DataSourceCreate, +) +from baserow_enterprise.assistant.tools.database.tools import ( + create_fields, + create_tables, + create_view_filters, + create_views, +) +from baserow_enterprise.assistant.tools.toolset import InlineRefsToolset + +from .utils import make_test_ctx + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "function, arguments", + [ + ( + create_tables, + {"database_id": 1, "tables": [], "add_sample_rows": False}, + ), + (create_fields, {"table_id": 1, "fields": []}), + (create_views, {"table_id": 1, "views": []}), + (create_view_filters, {"view_filters": []}), + ], +) +async def test_empty_creation_payload_retries_before_accessing_the_database( + function, arguments +): + """The missing database fixture also prevents unnoticed lookups or writes.""" + + model = TestModel() + toolset = InlineRefsToolset( + FunctionToolset([function]), model=model, model_profile=MagicMock() + ) + ctx = RunContext( + deps=make_test_ctx(None, None).deps, + model=model, + usage=RunUsage(), + prompt="Create items", + ) + tools = await toolset.get_tools(ctx) + + with pytest.raises(ModelRetry, match="Nothing was changed"): + await toolset.call_tool( + function.__name__, + {**arguments, "thought": "Creating items"}, + ctx, + tools[function.__name__], + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("table_id", [0, -1, "0"]) +async def test_placeholder_ids_are_rejected_before_repair_or_execution( + monkeypatch, table_id +): + executed = [] + + def read_table(table_id: int): + executed.append(table_id) + + model = TestModel() + toolset = InlineRefsToolset( + FunctionToolset([read_table]), model=model, model_profile=MagicMock() + ) + ctx = RunContext(deps=None, model=model, usage=RunUsage(), prompt="Read a table") + tools = await toolset.get_tools(ctx) + repair = AsyncMock() + monkeypatch.setattr(toolset, "_fix_tool_args", repair) + + result = await toolset.call_tool( + "read_table", {"table_id": table_id}, ctx, tools["read_table"] + ) + + assert "Not executed" in result["error"] + assert "list_tables" in result["next_steps"] + assert executed == [] + repair.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_missing_id_cannot_be_repaired_without_real_information(monkeypatch): + executed = [] + + def read_table(table_id: int): + executed.append(table_id) + + model = TestModel() + toolset = InlineRefsToolset( + FunctionToolset([read_table]), model=model, model_profile=MagicMock() + ) + ctx = RunContext(deps=None, model=model, usage=RunUsage(), prompt="Read a table") + tools = await toolset.get_tools(ctx) + repair = AsyncMock( + return_value=SimpleNamespace( + output=json.dumps({"__cannot_fix__": "table_id must come from list_tables"}) + ) + ) + monkeypatch.setattr( + "baserow_enterprise.assistant.tools.toolset.run_agent_with_model", repair + ) + + with pytest.raises(ModelRetry, match="table_id must come from list_tables"): + await toolset.call_tool("read_table", {}, ctx, tools["read_table"]) + + repair.assert_awaited_once() + assert executed == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("table_id", ["--1", "²"]) +async def test_malformed_ids_reach_argument_repair_without_running_the_tool( + monkeypatch, table_id +): + executed = [] + + def read_table(table_id: int): + executed.append(table_id) + + model = TestModel() + toolset = InlineRefsToolset( + FunctionToolset([read_table]), model=model, model_profile=MagicMock() + ) + ctx = RunContext(deps=None, model=model, usage=RunUsage(), prompt="Read a table") + tools = await toolset.get_tools(ctx) + repair = AsyncMock(side_effect=ModelRetry("Read the real table ID first.")) + monkeypatch.setattr(toolset, "_fix_tool_args", repair) + + with pytest.raises(ModelRetry, match="Read the real table ID first"): + await toolset.call_tool( + "read_table", {"table_id": table_id}, ctx, tools["read_table"] + ) + + repair.assert_awaited_once() + assert repair.call_args.args[:2] == ("read_table", {"table_id": table_id}) + assert executed == [] + + +@pytest.mark.parametrize("node_type", [[], {}], ids=["list", "object"]) +@pytest.mark.parametrize( + "payload_model, fields", + [ + pytest.param(TriggerNodeCreate, {"ref": "t", "label": "Trigger"}, id="trigger"), + pytest.param( + ActionNodeCreate, + {"ref": "a", "label": "Action", "previous_node_ref": "t"}, + id="action", + ), + pytest.param( + DataSourceCreate, + {"ref": "d", "name": "Source", "table_id": 1}, + id="data-source", + ), + ], +) +def test_malformed_type_aliases_produce_validation_errors( + payload_model, fields, node_type +): + with pytest.raises(ValidationError) as exc: + payload_model.model_validate({**fields, "type": node_type}) + + assert [(error["loc"], error["type"]) for error in exc.value.errors()] == [ + (("type",), "literal_error") + ] + + +@pytest.mark.asyncio +async def test_malformed_type_is_repaired_before_running_the_tool(monkeypatch): + executed = [] + + def read_source(data_source: DataSourceCreate): + executed.append(data_source) + return data_source.type + + model = TestModel() + toolset = InlineRefsToolset( + FunctionToolset([read_source]), model=model, model_profile=MagicMock() + ) + ctx = RunContext(deps=None, model=model, usage=RunUsage(), prompt="Read a source") + tools = await toolset.get_tools(ctx) + fields = {"ref": "d", "name": "Source", "table_id": 1} + repair = AsyncMock( + return_value=SimpleNamespace( + output=json.dumps( + {"data_source": {**fields, "type": "local_baserow_list_rows"}} + ) + ) + ) + monkeypatch.setattr( + "baserow_enterprise.assistant.tools.toolset.run_agent_with_model", repair + ) + + result = await toolset.call_tool( + "read_source", + {"data_source": {**fields, "type": []}}, + ctx, + tools["read_source"], + ) + + repair.assert_awaited_once() + assert result == "list_rows" + assert executed == [DataSourceCreate(**fields, type="list_rows")] diff --git a/enterprise/web-frontend/modules/baserow_enterprise/services/assistant.js b/enterprise/web-frontend/modules/baserow_enterprise/services/assistant.js index 75ef1dde15..0096c12071 100644 --- a/enterprise/web-frontend/modules/baserow_enterprise/services/assistant.js +++ b/enterprise/web-frontend/modules/baserow_enterprise/services/assistant.js @@ -1,3 +1,9 @@ +// json.dumps escapes newlines, so this can never occur inside a record. +const RECORD_DELIMITER = '\n\n' + +const STREAM_LOG_PREFIX = '[assistant stream]' +const CORRUPT_RECORD_PREVIEW_LENGTH = 200 + /** * The AI Assistant starts from the root URL, not the /api URL like the rest of * the Baserow API. This service file therefore overrides the baseURL to be the @@ -8,6 +14,60 @@ function getAssistantBaseURL(client) { return url.origin } +/** + * Turns the cumulative `xhr.responseText` into whole delimiter-terminated + * records and hands them to the consumer one at a time, in arrival order. + * + * @param {Function} readResponseText Returns everything received so far. + * @param {Function} onRecord Called with each parsed record, may be async. + * @returns {{read: Function, flush: Function}} Reader driven by the XHR events. + */ +function createRecordReader(readResponseText, onRecord) { + let consumed = 0 + let pending = '' + let dispatched = Promise.resolve() + + // Taking records synchronously keeps the buffer consistent across progress events. + const take = (isFinal) => { + const responseText = readResponseText() + pending += responseText.substring(consumed) + consumed = responseText.length + + const records = pending.split(RECORD_DELIMITER) + pending = isFinal ? '' : records.pop() + return records.filter((record) => record.trim() !== '') + } + + const dispatch = (records) => { + dispatched = dispatched.then(async () => { + for (const record of records) { + let update + try { + update = JSON.parse(record) + } catch (error) { + console.error(`${STREAM_LOG_PREFIX} dropped an unparsable record`, { + length: record.length, + preview: record.slice(0, CORRUPT_RECORD_PREVIEW_LENGTH), + error, + }) + continue + } + try { + await onRecord(update) + } catch (error) { + console.error(`${STREAM_LOG_PREFIX} record handler failed`, error) + } + } + }) + return dispatched + } + + return { + read: () => dispatch(take(false)), + flush: () => dispatch(take(true)), + } +} + export default (client) => { // Store active XHR request by chat UUID for cancellation const activeRequests = new Map() @@ -25,7 +85,10 @@ export default (client) => { adapter: (config) => { return new Promise((resolve, reject) => { const xhr = new XMLHttpRequest() - let buffer = '' + const reader = createRecordReader( + () => xhr.responseText, + onDownloadProgress ?? (() => {}) + ) // Store XHR for potential cancellation activeRequests.set(chatUuid, xhr) @@ -36,18 +99,7 @@ export default (client) => { }) xhr.onprogress = () => { - const chunk = xhr.responseText.substring(buffer.length) - buffer = xhr.responseText - - chunk.split('\n\n').forEach(async (line) => { - if (line.trim()) { - try { - await onDownloadProgress(JSON.parse(line)) - } catch (e) { - console.trace(e) - } - } - }) + reader.read() } xhr.onload = () => { @@ -56,7 +108,9 @@ export default (client) => { // Check if the request was successful (2xx status codes) if (xhr.status >= 200 && xhr.status < 300) { - resolve({ data: xhr.responseText, status: xhr.status }) + reader.flush().then(() => { + resolve({ data: xhr.responseText, status: xhr.status }) + }) } else { let errorData try { diff --git a/enterprise/web-frontend/test/unit/enterprise/services/assistant.spec.js b/enterprise/web-frontend/test/unit/enterprise/services/assistant.spec.js new file mode 100644 index 0000000000..1145b65acf --- /dev/null +++ b/enterprise/web-frontend/test/unit/enterprise/services/assistant.spec.js @@ -0,0 +1,288 @@ +import { afterEach, beforeEach, describe, expect, test, vi } from 'vitest' + +import AssistantService from '@baserow_enterprise/services/assistant' + +const CHAT_UUID = '11111111-1111-1111-1111-111111111111' + +const EVENTS = [ + { type: 'ai_started', message_id: 7 }, + { type: 'thinking', content: 'Looking at your table' }, + // Embedded newlines prove the delimiter survives json.dumps escaping. + { type: 'reasoning', id: 'r1', content: 'first\n\nsecond' }, + { type: 'message', id: 'm1', content: 'Ünïcödé 😀 answer', sources: [] }, +] + +const BODY = EVENTS.map((event) => JSON.stringify(event) + '\n\n').join('') + +class FakeXHR { + constructor() { + this.responseText = '' + this.status = 200 + this.statusText = 'OK' + this.aborted = false + FakeXHR.instances.push(this) + } + + open() {} + + setRequestHeader() {} + + send() {} + + abort() { + this.aborted = true + this.onabort() + } + + receive(text) { + this.responseText += text + this.onprogress() + } + + complete(status = 200) { + this.status = status + this.onload() + } +} + +FakeXHR.instances = [] + +const makeClient = () => ({ + defaults: { baseURL: 'http://localhost:8000/api' }, + post: vi.fn((url, data, config) => + config.adapter({ + baseURL: config.baseURL, + url, + headers: { 'Content-Type': 'application/json' }, + data: JSON.stringify(data), + }) + ), +}) + +/** + * Streams `body` through the adapter using the given chunk boundaries and + * returns every record the consumer received. + */ +const streamBody = async (body, boundaries, { complete = true } = {}) => { + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + (update) => { + received.push(update) + } + ) + const xhr = FakeXHR.instances[FakeXHR.instances.length - 1] + + let offset = 0 + for (const boundary of [...boundaries, body.length]) { + if (boundary > offset) { + xhr.receive(body.substring(offset, boundary)) + offset = boundary + } + } + if (complete) { + xhr.complete() + await request + } + return { received, request, xhr } +} + +describe('assistant service streaming adapter', () => { + beforeEach(() => { + FakeXHR.instances = [] + vi.stubGlobal('XMLHttpRequest', FakeXHR) + }) + + afterEach(() => { + vi.unstubAllGlobals() + vi.restoreAllMocks() + }) + + test('delivers every record exactly once for every possible split', async () => { + for (let split = 0; split <= BODY.length; split++) { + const { received } = await streamBody(BODY, [split]) + expect({ split, received }).toEqual({ split, received: EVENTS }) + } + }) + + test('delivers every record exactly once for every three-way split', async () => { + for (let first = 0; first <= BODY.length; first += 7) { + for (let second = first; second <= BODY.length; second += 11) { + const { received } = await streamBody(BODY, [first, second]) + expect({ first, second, received }).toEqual({ + first, + second, + received: EVENTS, + }) + } + } + }) + + test('delivers every record when a chunk ends inside the delimiter', async () => { + const delimiterOffsets = [] + for (let i = 0; i < BODY.length - 1; i++) { + if (BODY.substring(i, i + 2) === '\n\n') { + delimiterOffsets.push(i, i + 1, i + 2) + } + } + expect(delimiterOffsets.length).toBe(EVENTS.length * 3) + + for (const offset of delimiterOffsets) { + const { received } = await streamBody(BODY, [offset]) + expect({ offset, received }).toEqual({ offset, received: EVENTS }) + } + }) + + // responseText never exposes a partial multi-byte sequence, so this is the worst case. + test('delivers a record whose astral character is split across chunks', async () => { + const emojiIndex = BODY.indexOf('😀') + expect(BODY.charCodeAt(emojiIndex)).toBeGreaterThanOrEqual(0xd800) + + const { received } = await streamBody(BODY, [emojiIndex + 1]) + expect(received).toEqual(EVENTS) + }) + + test('delivers several complete records from a single progress event', async () => { + const { received } = await streamBody(BODY, []) + expect(received).toEqual(EVENTS) + }) + + test('does not redeliver records when a progress event adds no bytes', async () => { + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + (update) => received.push(update) + ) + const xhr = FakeXHR.instances[0] + + xhr.receive(BODY) + xhr.onprogress() + xhr.onprogress() + xhr.complete() + await request + + expect(received).toEqual(EVENTS) + }) + + test('flushes a final record that never received its delimiter', async () => { + const truncated = BODY.slice(0, -'\n\n'.length) + const { received } = await streamBody(truncated, [truncated.length - 5]) + + expect(received).toEqual(EVENTS) + }) + + test('reports a truncated tail instead of delivering it', async () => { + const error = vi.spyOn(console, 'error').mockImplementation(() => {}) + const truncated = BODY + '{"type": "message", "content": "cut' + const { received } = await streamBody(truncated, [BODY.length]) + + expect(received).toEqual(EVENTS) + expect(error).toHaveBeenCalledTimes(1) + expect(error.mock.calls[0][0]).toContain('unparsable record') + }) + + test('skips a corrupt record but keeps delivering the rest', async () => { + const error = vi.spyOn(console, 'error').mockImplementation(() => {}) + const corrupt = + JSON.stringify(EVENTS[0]) + + '\n\nNone' + + JSON.stringify(EVENTS[1]) + + '\n\n' + + JSON.stringify(EVENTS[2]) + + '\n\n' + const { received } = await streamBody(corrupt, []) + + expect(received).toEqual([EVENTS[0], EVENTS[2]]) + expect(error).toHaveBeenCalledTimes(1) + expect(error.mock.calls[0][0]).toContain('unparsable record') + }) + + test('keeps record order when the consumer is asynchronous', async () => { + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + async (update) => { + await new Promise((resolve) => setTimeout(resolve, 0)) + received.push(update) + } + ) + const xhr = FakeXHR.instances[0] + + for (const event of EVENTS) { + xhr.receive(JSON.stringify(event) + '\n\n') + } + xhr.complete() + await request + + expect(received).toEqual(EVENTS) + }) + + test('resolves only after every record has been handled', async () => { + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + async (update) => { + await new Promise((resolve) => setTimeout(resolve, 0)) + received.push(update) + } + ) + const xhr = FakeXHR.instances[0] + + xhr.receive(BODY) + xhr.complete() + await request + + expect(received).toEqual(EVENTS) + }) + + test('a throwing consumer does not stop the remaining records', async () => { + const error = vi.spyOn(console, 'error').mockImplementation(() => {}) + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + (update) => { + if (update.type === 'thinking') { + throw new Error('consumer exploded') + } + received.push(update) + } + ) + const xhr = FakeXHR.instances[0] + + xhr.receive(BODY) + xhr.complete() + await request + + expect(received).toEqual([EVENTS[0], EVENTS[2], EVENTS[3]]) + expect(error.mock.calls[0][0]).toContain('record handler failed') + }) + + test('rejects without delivering records on a non 2xx response', async () => { + const received = [] + const request = AssistantService(makeClient()).sendMessage( + CHAT_UUID, + 'hello', + {}, + (update) => received.push(update) + ) + const xhr = FakeXHR.instances[0] + + xhr.responseText = JSON.stringify({ detail: 'nope', error: 'ERROR_X' }) + xhr.complete(400) + + await expect(request).rejects.toMatchObject({ + response: { status: 400, data: { error: 'ERROR_X' } }, + }) + expect(received).toEqual([]) + }) +}) diff --git a/premium/backend/src/baserow_premium/prompts/__init__.py b/premium/backend/src/baserow_premium/prompts/__init__.py index f410af8a77..fdde8008ff 100644 --- a/premium/backend/src/baserow_premium/prompts/__init__.py +++ b/premium/backend/src/baserow_premium/prompts/__init__.py @@ -30,4 +30,5 @@ def get_generate_formula_prompt(): @cache def get_formula_docs(): + # Callers embed this in a str.format() template, so it must stay brace-free. return read_text("baserow_premium.prompts", "formula_docs.md") diff --git a/premium/backend/src/baserow_premium/prompts/formula_docs.md b/premium/backend/src/baserow_premium/prompts/formula_docs.md index 50a73de63d..9c0da17be0 100644 --- a/premium/backend/src/baserow_premium/prompts/formula_docs.md +++ b/premium/backend/src/baserow_premium/prompts/formula_docs.md @@ -63,6 +63,58 @@ To create formulas to make a Boolean test on data in field C, taking data from f Using `join()` to convert the list to text handles the empty scenario correctly. This formula checks if the Organization field (a link-to-table field) has a value. If it's true, it shows the content of the Name field; otherwise, it displays the content of the Notes field. +## Hard Rules + +These are the constraints most often got wrong. Breaking any of them makes the +formula fail to compile. + +1. **Reference every field with `field('Name')`.** Curly-brace references + borrowed from other tools are not part of this language and cannot even be + tokenized, so the parser rejects them before any type checking happens. + Write `field('Total')`, never a brace-wrapped field name. +2. **`and()` and `or()` take exactly 2 arguments.** For three or more conditions, + nest them: `or(a, or(b, c))`, `and(a, and(b, c))`. Passing 3+ arguments fails + with "N arguments were given to the 'or' function". +3. **`min()` and `max()` take a single array argument**, such as a `lookup()` or a + link/lookup `field()`. They are not two-argument scalar functions. For the + larger of two values use `if(a > b, a, b)`. +4. **`datetime_format()` format strings are PostgreSQL `to_char` patterns**, not + moment.js or Excel patterns. Use `'Day'` (or `'FMDay'` unpadded) for a weekday + name, `'Month'` for a month name, `'YYYY-MM-DD'` for a date, `'HH24:MI'` for a + 24-hour time. Lowercase moment tokens such as `'dddd'`, `'mmmm'` or `'hh:mm'` + are silently wrong and produce numbers rather than names. +5. **Only use functions listed in the Function Reference below.** Baserow is not + Excel: `weekday`, `substitute`, `countif`, `iferror`, `datedif` and `sumif` do + not exist. Their nearest equivalents here are `todate`/`datetime_format`, + `replace`/`regex_replace`, `sum(filter(...))`, `when_empty` and `date_diff`. + +## Type Coercion + +Most functions accept only certain argument types. When the value you have is not the +type a function wants, convert it: a type mismatch is a reason to add a conversion, never +a reason to conclude that something cannot be expressed. The error names the argument +position and the type it received, and lists the accepted types when there are any. + +| You have | You want | Use | Notes | +| -------- | -------- | --- | ----- | +| any valid type | text | `totext(x)` | `totext` accepts every valid type, including single select, multiple select, multiple collaborators, link row, lookup, date, number and boolean. | +| single select | text | `totext(field('x'))` | Returns the selected option's value, or `''` when the cell is empty. | +| multiple select / multiple collaborators | text | `totext(field('x'))` | Returns the values joined into one string. The same holds for a `lookup()` of such a field. | +| a list (link row, lookup, array) | text | `join(field('x'), ', ')` | `join` takes a list plus a separator. If the items are not text, convert them first: `join(totext(field('x')), ', ')`. | +| text | number | `tonumber(x)` | `tonumber` accepts text only; wrap anything else: `tonumber(totext(field('x')))`. | +| any value | a non-empty value | `when_empty(x, fallback)` | Both arguments must be the same type. | +| a list (link row, lookup, multiple select) | item count | `count(field('x'))` | | +| multiple select | test for an option | `has_option(field('x'), 'value')` | | + +`concat(...)` applies `totext` to every argument itself, so mixing types inside it needs +no wrapping. + +`=` and `!=` cast both sides to text when the two types differ but are comparable, so a +single select can be compared with text directly. A multiple select is comparable only +with another multiple select: test its contents with `has_option`, or compare +`totext(field('x'))`. The ordering operators `>`, `>=`, `<` and `<=` refuse select fields +in either position. + ## Function Reference ### Text Functions @@ -163,7 +215,24 @@ The `today()` function is useful for calculating intervals or when you need to h ### Aggregate Functions -These functions work with arrays and lookup values to perform calculations across multiple values. +These functions collapse a **list** of values into one value. A list only ever comes +from a reference to a link row field, a reference to a lookup field, or a `lookup()` +call. A plain field of the current row holds a single value and can never be +aggregated — combine values within one row with the arithmetic operators, and use an +aggregate function only to combine values across linked rows. + +A `field()` reference to a link row field carries the **type of the linked table's +primary field**, as a list; `lookup()` carries the type of the field it names. Arguments +are type checked and never coerced, so that type must already be one the function +accepts: `sum` and `avg` take numbers or durations; `min` and `max` also take text and +dates; `every` and `any` take booleans; `join` takes a lookup field reference or a list +of text; `count` takes a list of any type, and also a multiple select or multiple +collaborators field directly. + +When a list is not yet the type a function needs, convert it rather than concluding the +conversion is unsupported. `totext()` accepts every valid type, including multiple +select and multiple collaborators, so `join(totext(field('a link row field')), ', ')` is +the general way to render any list as text. | Functions | Details | Syntax | Examples | | --------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | ------------------------------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | diff --git a/premium/backend/tests/baserow_premium_tests/prompts/test_prompt_assets.py b/premium/backend/tests/baserow_premium_tests/prompts/test_prompt_assets.py new file mode 100644 index 0000000000..55ddb2b682 --- /dev/null +++ b/premium/backend/tests/baserow_premium_tests/prompts/test_prompt_assets.py @@ -0,0 +1,152 @@ +import ast +import re +from pathlib import Path + +import pytest + +from baserow.contrib.database.formula.registries import formula_function_registry +from baserow_premium.fields import handler as ai_field_handler +from baserow_premium.prompts import get_formula_docs, get_generate_formula_prompt + +PROMPT_BUILDER = "get_generate_formula_prompt" +FUNCTION_REFERENCE_HEADING = "## Function Reference" +INLINE_CODE = re.compile(r"`([^`\n]+)`") +FUNCTION_CALL = re.compile(r"\b([a-z_][a-z0-9_]*)\s*\(") + +# Parsed straight from the grammar as references, so they never reach the registry. +GRAMMAR_FUNCTIONS = frozenset({"field", "lookup"}) + + +def _handler_format_keys() -> set[str]: + """The keyword names the AI field handler passes when formatting the prompt.""" + + handler_path = Path(ai_field_handler.__file__) + for node in ast.walk(ast.parse(handler_path.read_text())): + if not isinstance(node, ast.Call) or not isinstance(node.func, ast.Attribute): + continue + builder = node.func.value + if ( + node.func.attr == "format" + and isinstance(builder, ast.Call) + and getattr(builder.func, "id", None) == PROMPT_BUILDER + ): + return {keyword.arg for keyword in node.keywords if keyword.arg is not None} + + pytest.fail( + f"No `{PROMPT_BUILDER}().format(...)` call found in {handler_path}. Either the " + "prompt is no longer formatted there, in which case this guard must follow the " + "new consumer, or it is formatted with **kwargs and the keys can no longer be " + "derived statically." + ) + + +def _registered_function_names() -> set[str]: + return set(formula_function_registry.registry.keys()) | GRAMMAR_FUNCTIONS + + +def _documented_function_names(docs: str) -> set[str]: + """The names in the first column of every Function Reference table.""" + + _, heading, reference = docs.partition(FUNCTION_REFERENCE_HEADING) + assert heading, ( + f"'{FUNCTION_REFERENCE_HEADING}' heading is gone from formula_docs.md, so this " + "guard can no longer find the documented functions." + ) + + names = set() + for line in reference.splitlines(): + if not line.startswith("|"): + continue + cell = line.split("|")[1].strip() + if not cell or cell == "Functions" or set(cell) <= set("-: "): + continue + # Operator rows read "add `+`"; the function name is what precedes the operator. + names.add(cell.split("`")[0].strip()) + return names + + +def _function_names_called_in_docs(docs: str) -> set[str]: + """The names used in call form inside inline code, such as the Hard Rules examples. + + Only call form counts. The doc also names Excel functions like `countif` in prose, + precisely to say they do not exist in Baserow. + """ + + return { + match.group(1) + for span in INLINE_CODE.findall(docs) + for match in FUNCTION_CALL.finditer(span) + } + + +def test_generate_formula_prompt_formats_with_the_keys_the_handler_passes(): + keys = _handler_format_keys() + assert keys, f"{PROMPT_BUILDER}() is formatted without any keyword argument." + + values = {key: f"<<{key} sentinel>>" for key in keys} + try: + formatted = get_generate_formula_prompt().format(**values) + except (KeyError, IndexError, ValueError) as exc: + pytest.fail( + f"{PROMPT_BUILDER}().format({', '.join(sorted(keys))}) raised " + f"{type(exc).__name__}: {exc}. The prompt is the formula docs concatenated " + "with INSTRUCTIONS, so a brace introduced anywhere in either piece is read " + "as a placeholder and breaks AI formula generation at runtime." + ) + + for key, sentinel in sorted(values.items()): + assert sentinel in formatted, ( + f"The prompt no longer substitutes '{key}'. The handler still passes it, so " + "that data never reaches the model." + ) + + +def test_formula_docs_contain_no_literal_brace(): + docs = get_formula_docs() + offenders = [ + f"line {number}: {line.strip()}" + for number, line in enumerate(docs.splitlines(), start=1) + if "{" in line or "}" in line + ] + + assert not offenders, ( + "formula_docs.md must contain no literal brace: " + f"{PROMPT_BUILDER}() concatenates it into a str.format() template that " + f"{Path(ai_field_handler.__file__).name} formats, so every brace is read as a " + "placeholder and an unknown one raises KeyError for every AI formula " + "generation. Do not escape them as " + "{{ }} either: the assistant reads the same file unformatted through " + "get_formula_docs(), and the model would then see the doubled braces. Reword " + "the text instead. Offending lines:\n" + "\n".join(offenders) + ) + + +def test_function_reference_only_documents_registered_functions(): + documented = _documented_function_names(get_formula_docs()) + assert len(documented) > 50, ( + f"Only {len(documented)} functions were parsed out of the Function Reference " + "tables, so their layout changed and this guard checks almost nothing." + ) + + unknown = sorted(documented - _registered_function_names()) + assert not unknown, ( + f"The Function Reference documents {unknown}, which are not registered in " + "backend/src/baserow/contrib/database/formula/ast/function_defs.py. A document " + "whose job is to stop the model inventing functions must not invent any itself." + ) + + +def test_functions_used_in_doc_examples_are_registered(): + called = _function_names_called_in_docs(get_formula_docs()) + assert len(called) > 10, ( + f"Only {len(called)} called functions were parsed out of formula_docs.md, so " + "the inline code examples changed shape and this guard checks almost nothing." + ) + + unknown = sorted(called - _registered_function_names()) + assert not unknown, ( + f"formula_docs.md shows {unknown} being called in its rules and examples, but " + "they are not registered in " + "backend/src/baserow/contrib/database/formula/ast/function_defs.py. Every " + "formula the doc demonstrates must actually compile." + ) diff --git a/web-frontend/yarn.lock b/web-frontend/yarn.lock index df7c81b6c3..20bfb01767 100644 --- a/web-frontend/yarn.lock +++ b/web-frontend/yarn.lock @@ -6248,9 +6248,9 @@ detect-libc@^2.0.0, detect-libc@^2.0.3: integrity sha512-Btj2BOOO83o3WyH59e8MgXsxEQVcarkUOpEYrubB0urwnN10yQ364rsiByU11nZlqWYZm05i/of7io4mzihBtQ== devalue@^5.8.2: - version "5.9.0" - resolved "https://registry.yarnpkg.com/devalue/-/devalue-5.9.0.tgz#5d30db41a0db9171cf4ee9dbf6617e0b77cd7c2a" - integrity sha512-RWrqdArjvPbsATEhOPUo6Wndc/iWnkWKlhIrdlF3zMMYo/c3CVtoaVAyLtWxz5h8nSlkHzxnzV2uLydPXmtF+A== + version "5.9.2" + resolved "https://registry.yarnpkg.com/devalue/-/devalue-5.9.2.tgz#2a3a8ad21904c6a630bf7bb160acb7ca2fa1e467" + integrity sha512-po4PAY5c53tw5XMocSnf8A/5OHhbbUftpr93aEN6BBoAdntUmK7vu7wOATqvt7cXO7m1Cl4gMVn6p7n6n4mj0w== devframe@^0.5.2: version "0.5.4"