sadpandajoe commented on code in PR #43134:
URL: https://github.com/apache/superset/pull/43134#discussion_r3871832690
##########
superset/config.py:
##########
@@ -2948,6 +2964,326 @@ def EMAIL_HEADER_MUTATOR( # pylint:
disable=invalid-name,unused-argument # noq
"CACHE_REDIS_SSL_CA_CERTS": None,
}
+# ---------------------------------------------------------
+# AI assistant
+# ---------------------------------------------------------
+# Requires the AI_ASSISTANT feature flag. Superset ships no model provider and
+# talks to no model vendor by default: until AI_LLM_PROVIDER_CLASS names a
+# usable provider the assistant's endpoints return 404.
+#
+# Dotted path to a superset.ai.llm.base.BaseLLMProvider subclass. Point this at
+# a vendor provider, an OpenAI-compatible endpoint, a self-hosted model, or a
+# private gateway. Everything vendor-specific — base URLs, authentication,
+# model naming — belongs in the provider, not here.
+AI_LLM_PROVIDER_CLASS: str | None = None
+
+# Keyword arguments passed to the provider's constructor. Contents are entirely
+# provider-defined. Keep credentials out of this file: read them from the
+# environment or a secret store in your own config.
+#
+# AI_LLM_PROVIDER_CONFIG = {
+# "api_key": os.environ["MY_LLM_API_KEY"],
+# "base_url": "https://llm.internal.example.com/v1",
+# "models": {
+# "default": "some-balanced-model",
+# "fast": "some-small-model",
+# "reasoning": "some-large-model",
+# },
+# }
+AI_LLM_PROVIDER_CONFIG: dict[str, Any] = {}
+
+# Dotted path to a superset.ai.runtime.base.BaseAgentRuntime subclass driving
+# the tool-use loop.
+AI_AGENT_RUNTIME_CLASS = "superset.ai.runtime.messages.MessagesApiRuntime"
+
+# Where a turn is executed.
+#
+# "inline" — in the web worker handling the request. No extra
infrastructure,
+# but a turn occupies a worker for its whole duration.
+# "worker" — handed to Celery; the request streams events from the event
bus.
+# Survives a browser reconnect and keeps web workers free, at the
+# cost of requiring Celery and a shared event bus.
+AI_ASSISTANT_EXECUTION_MODE: Literal["inline", "worker"] = "inline"
+
+# How streamed events travel from producer to the HTTP response.
+#
+# "memory" — an in-process queue. Correct only when the producer and the
+# streaming request are the same process, i.e. inline execution.
+# "redis" — Redis streams, via the same cache backend the async-query
+# channel uses. Required for "worker" execution mode.
+AI_ASSISTANT_EVENT_BUS: Literal["memory", "redis"] = "memory"
+
+# Redis connection for the AI event bus. Required when AI_ASSISTANT_EVENT_BUS
is
+# "redis". Streams need commands the general-purpose cache client does not
+# expose, so this is configured separately rather than borrowed from
+# CACHE_CONFIG. The accepted shape matches
+# GLOBAL_ASYNC_QUERIES_CACHE_BACKEND; point both at the same Redis if you like.
+AI_ASSISTANT_EVENT_BUS_CACHE_CONFIG: dict[str, Any] = {
+ "CACHE_TYPE": "RedisCache",
+ "CACHE_REDIS_HOST": "localhost",
+ "CACHE_REDIS_PORT": 6379,
+ "CACHE_REDIS_USER": "",
+ "CACHE_REDIS_PASSWORD": "",
+ "CACHE_REDIS_DB": 0,
+ "CACHE_DEFAULT_TIMEOUT": 300,
+ "CACHE_REDIS_SSL": False,
+}
+
+# Key prefix for AI event streams when the Redis bus is in use.
+AI_ASSISTANT_EVENT_STREAM_PREFIX = "ai-events-"
+
+# How long a run's event stream is retained, in seconds. Bounds how late a
+# reconnecting browser can still pick up a run it lost.
+AI_ASSISTANT_EVENT_TTL_SECONDS = 900
+
+# Named agent profiles, merged over the built-ins by key. Each value is a dict
+# of fields to override, so narrowing one profile does not mean restating the
+# rest. The most important field is "tools": which tools that profile may
+# invoke. An unknown tool name is a startup error, not a silent omission.
+#
+# AI_AGENT_PROFILES = {
+# # Take the shipped default but forbid raw SQL.
+# "default": {"tools": ["search_assets", "get_schema"]},
+# # Let the analyst profile think harder and longer.
+# "analyst": {"model_alias": "reasoning", "max_turns": 60},
+# # Add a profile only some users may select.
+# "deep": {
+# "name": "Deep analysis",
+# "tools": ["search_assets", "get_schema", "execute_sql"],
+# "required_permission": ("can_write", "AIAssistant"),
+# },
+# }
+AI_AGENT_PROFILES: dict[str, Any] = {}
+
+# Ceiling on model round trips in a single turn. A turn that needs more than
+# this is answered with what it has rather than looping indefinitely.
+AI_AGENT_MAX_TURNS = 20
+
+# Wall-clock budget for one turn, in seconds.
+AI_AGENT_TIMEOUT_SECONDS = 300
+
+# Pre-tool-use guards, applied in order. Each is a dotted path to a
+# superset.ai.policy.ToolPolicy implementation. These bound blast radius; they
+# do not replace the per-object authorization checks inside each tool.
+AI_AGENT_TOOL_POLICIES: list[str] = [
+ "superset.ai.policy.ReadOnlySqlPolicy",
+ "superset.ai.policy.IdentifierPolicy",
+ "superset.ai.policy.ForeignToolPolicy",
+]
+
+# Rows and bytes a single tool result may return before it is truncated.
+# Model context is finite, and an unbounded result set exhausts it.
+AI_AGENT_MAX_RESULT_ROWS = 500
+AI_AGENT_MAX_RESULT_BYTES = 256 * 1024
+
+# External MCP servers whose tools may be offered to an agent profile. Superset
+# ships none and integrates with no third-party service: with this empty,
nothing
+# in superset.ai.mcp is ever reached and the assistant behaves exactly as it
does
+# without it.
+#
+# A server listed here is only *available*. It is used by an agent profile that
+# names it in its "mcp_servers" field, via AI_AGENT_PROFILES. A profile naming
a
+# server that is not configured here is an error, not a silently shorter tool
+# list.
+#
+# AI_AGENT_MCP_SERVERS = {
+# # The key is the server name. It becomes part of every tool name this
+# # server contributes, so keep it short: letters, digits, hyphens and
+# # underscores, and no double underscore.
+# "acme_catalog": {
+# # Required. Absolute http:// or https:// endpoint.
+# "url": "https://mcp.acme.internal/mcp",
+# # "streamable_http" (default) or "sse".
+# "transport": "streamable_http",
+# # The ONLY headers sent to this server. Superset never forwards the
+# # user's session cookie, CSRF token or any Superset auth header: an
+# # external server is not a party to the user's Superset session.
+# # Read secrets from the environment rather than writing them here.
+# "headers": {"Authorization": f"Bearer
{os.environ['ACME_MCP_TOKEN']}"},
+# # Per-call budget. Bounds how long one call may occupy the worker
+# # running the turn. Defaults to 30.
+# "timeout_seconds": 30,
+# # Which of the server's tools to take. Absent or None means every
+# # tool it offers, which lets the server decide what the agent can
do.
+# # Either the server's own name ("search_tables") or the namespaced
+# # name Superset assigns ("mcp__acme_catalog__search_tables")
matches.
+# "tool_allowlist": ["search_tables"],
+# # Refused regardless of the allowlist.
+# "tool_denylist": [],
+# },
+# }
+#
+# AI_AGENT_PROFILES = {
+# "default": {"mcp_servers": ["acme_catalog"]},
+# }
+#
+# Every tool from a server is namespaced "mcp__<server>__<tool>". The
namespace is
+# stable, appears in stored conversation history, and is what makes it
impossible
+# for a server offering "execute_sql" to displace Superset's own tool of that
+# name. Foreign results pass through the same AI_AGENT_MAX_RESULT_BYTES bound
and
+# the same AI_AGENT_TOOL_POLICIES chain as built-in ones, and are wrapped as
+# untrusted content before the model sees them.
+#
+# A server that is unreachable, slow or unreadable contributes no tools and the
+# agent keeps working with the built-ins. Discovery happens while assembling
the
+# registry for a turn, so a slow server costs up to its timeout at the start of
+# each turn that uses it.
+#
+# Requires the 'mcp' package; it is imported only once a server is configured.
+AI_AGENT_MCP_SERVERS: dict[str, Any] = {}
+
+# Refuse any external MCP tool whose name advertises SQL execution — anything
+# containing "execute_sql", "run_sql" or "query" by default. Enforced by
+# superset.ai.policy.ForeignToolPolicy.
+#
+# On by default because Superset's read-only enforcement and its per-datasource
+# authorization can only apply to SQL Superset itself runs. A third-party
server
+# executing SQL goes through neither, so permitting it silently removes both
+# controls rather than merely widening the surface. Set this False only if you
+# have satisfied yourself that the servers you have configured enforce
+# equivalent controls of their own.
+AI_AGENT_MCP_DENY_FOREIGN_SQL = True
+
+# Conversation history sent to the model: the most recent N messages, further
+# trimmed oldest-first until under the character budget.
+AI_ASSISTANT_MAX_HISTORY_MESSAGES = 25
+AI_ASSISTANT_MAX_HISTORY_CHARS = 100_000
+
+# Timezone for the authoritative date given to the model, so it never has to
+# infer today's date or weekday.
+AI_ASSISTANT_TIMEZONE = "UTC"
+
+# Days a conversation is retained. Pruning is performed by the
+# ``ai.prune_conversations`` Celery task, which must be scheduled to run.
+AI_ASSISTANT_MESSAGE_RETENTION_DAYS = 30
Review Comment:
This advertises retention through `ai.prune_conversations`, but no such task
or consumer of `AI_ASSISTANT_MESSAGE_RETENTION_DAYS` exists in this change.
Operators relying on the default will retain conversations and stored tool
details indefinitely. Could the pruning task and schedule hook be included
before documenting this setting?
##########
superset/ai/api.py:
##########
@@ -0,0 +1,996 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""
+REST API for the AI assistant.
+
+Every route carries ``@protect()`` and is reached through ``@expose`` on a
+``BaseSupersetApi`` subclass, which is what makes Flask-AppBuilder's
+authorization actually run. Ownership is enforced a second time in the command
+and DAO layers, so a conversation identifier is never on its own a capability.
+"""
+
+from __future__ import annotations
+
+import logging
+import time
+from collections.abc import Generator
+from typing import Any, cast
+
+from flask import current_app, request, Response, stream_with_context
+from flask_appbuilder.api import expose, permission_name, protect, safe
+from marshmallow import ValidationError
+
+from superset.ai.events import (
+ error_event,
+ KEEPALIVE_FRAME,
+ KEEPALIVE_INTERVAL_SECONDS,
+)
+from superset.ai.schemas import (
+ AgentResponseSchema,
+ CancelPostSchema,
+ FeedbackPostSchema,
+ MessagePostSchema,
+ RunAcceptedResponseSchema,
+ SuggestedPromptsPostSchema,
+ ThreadDetailResponseSchema,
+ ThreadPostSchema,
+ ThreadPutSchema,
+ ThreadResponseSchema,
+)
+from superset.ai.types import MessageRole, MessageStatus
+from superset.commands.ai.exceptions import (
+ AIChatMessageInvalidError,
+ AIChatMessageNotFoundError,
+ AIChatThreadInvalidError,
+ AIChatThreadNotFoundError,
+)
+from superset.extensions import event_logger
+from superset.utils.core import get_user_id
+from superset.utils.decorators import transaction
+from superset.views.base_api import BaseSupersetApi, statsd_metrics
+
+logger = logging.getLogger(__name__)
+
+#: Upper bound on how long a client may hold a stream open, so an abandoned
+#: browser tab cannot pin a worker indefinitely.
+_STREAM_TIMEOUT_SECONDS = 900
+
+#: How often a reader checks the event bus for new frames.
+#:
+#: Deliberately separate from ``KEEPALIVE_INTERVAL_SECONDS``. Passing the
+#: keep-alive interval as the poll interval made the reader sleep fifteen
seconds
+#: between checks and then deliver everything that had accumulated in one
batch —
+#: so a worker-mode run showed no streaming at all: the answer and every tool
call
+#: appeared in fifteen-second lumps. One controls responsiveness, the other how
+#: often an idle connection is reassured; they are not the same number.
+_EVENT_POLL_SECONDS = 0.1
+
+
+class AIRestApi(BaseSupersetApi):
+ """Conversations with the AI assistant."""
+
+ resource_name = "ai"
+ openapi_spec_tag = "AI Assistant"
+ allow_browser_login = True
+ class_permission_name = "AIAssistant"
+
+ openapi_spec_component_schemas = (
+ AgentResponseSchema,
+ CancelPostSchema,
+ FeedbackPostSchema,
+ MessagePostSchema,
+ RunAcceptedResponseSchema,
+ SuggestedPromptsPostSchema,
+ ThreadDetailResponseSchema,
+ ThreadPostSchema,
+ ThreadPutSchema,
+ ThreadResponseSchema,
+ )
+
+ @expose("/agent/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def agents(self) -> Response:
+ """List agent profiles the current user may select.
+ ---
+ get:
+ summary: List available agent profiles
+ responses:
+ 200:
+ description: Available profiles
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ type: array
+ items:
+ $ref: '#/components/schemas/AgentResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 403:
+ $ref: '#/components/responses/403'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.factories import get_profiles
+
+ profiles = get_profiles().visible_to_current_user()
+ return self.response(200, result=[p.to_public_dict() for p in
profiles])
+
+ @expose("/model/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def models(self) -> Response:
+ """List models this deployment has configured.
+ ---
+ get:
+ summary: List selectable models
+ responses:
+ 200:
+ description: Configured model identifiers
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ type: array
+ items:
+ type: string
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.factories import get_provider
+
+ return self.response(200, result=get_provider().available_models())
+
+ @expose("/thread/", methods=("POST",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.post_thread",
+ log_to_statsd=False,
+ )
+ def post_thread(self) -> Response:
+ """Create a conversation.
+ ---
+ post:
+ summary: Create a conversation
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/ThreadPostSchema'
+ responses:
+ 201:
+ description: Conversation created
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/ThreadResponseSchema'
+ 400:
+ $ref: '#/components/responses/400'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import CreateAIChatThreadCommand
+
+ try:
+ payload = ThreadPostSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ try:
+ thread = CreateAIChatThreadCommand(
+ user_id=self._user_id(),
+ title=payload.get("title"),
+ agent_key=payload.get("agent_key"),
+ ).run()
+ except AIChatThreadInvalidError as ex:
+ return self.response_422(message=str(ex))
+ return self.response(201, result=_thread_dict(thread))
+
+ @expose("/thread/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def get_threads(self) -> Response:
+ """List the current user's conversations.
+ ---
+ get:
+ summary: List conversations
+ parameters:
+ - in: query
+ name: limit
+ schema:
+ type: integer
+ - in: query
+ name: offset
+ schema:
+ type: integer
+ responses:
+ 200:
+ description: Conversations
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ count:
+ type: integer
+ result:
+ type: array
+ items:
+ $ref: '#/components/schemas/ThreadResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import AIChatThreadDAO
+
+ limit = request.args.get("limit", type=int) or 50
+ offset = request.args.get("offset", type=int) or 0
+ threads = AIChatThreadDAO.find_all_for_user(
+ self._user_id(), limit=limit, offset=offset
+ )
+ return self.response(
+ 200,
+ count=len(threads),
+ result=[_thread_dict(thread) for thread in threads],
+ )
+
+ @expose("/thread/<thread_uuid>", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def get_thread(self, thread_uuid: str) -> Response:
+ """Fetch a conversation and its messages.
+ ---
+ get:
+ summary: Get a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ responses:
+ 200:
+ description: Conversation with messages
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/ThreadDetailResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import (
+ AIChatFeedbackDAO,
+ AIChatMessageDAO,
+ AIChatThreadDAO,
+ )
+
+ user_id = self._user_id()
+ thread = AIChatThreadDAO.find_by_uuid_for_user(thread_uuid, user_id)
+ if thread is None:
+ return self.response_404()
+
+ messages = AIChatMessageDAO.find_for_thread(thread)
+ # Resolved for the whole transcript at once so the panel can show which
+ # replies this user already rated; without it a reload loses the
verdict
+ # and the message looks unrated.
+ verdicts = AIChatFeedbackDAO.find_verdicts_for_user(
+ [message.id for message in messages], user_id
+ )
+ detail = _thread_dict(thread)
+ detail["messages"] = [
+ _message_dict(message, liked=verdicts.get(message.id))
+ for message in messages
+ ]
+ return self.response(200, result=detail)
+
+ @expose("/thread/<thread_uuid>", methods=("PUT",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.put_thread",
+ log_to_statsd=False,
+ )
+ def put_thread(self, thread_uuid: str) -> Response:
+ """Rename or archive a conversation.
+ ---
+ put:
+ summary: Update a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/ThreadPutSchema'
+ responses:
+ 200:
+ description: Conversation updated
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ 422:
+ $ref: '#/components/responses/422'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import UpdateAIChatThreadCommand
+
+ try:
+ payload = ThreadPutSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ try:
+ thread = UpdateAIChatThreadCommand(
+ thread_uuid,
+ self._user_id(),
+ title=payload.get("title"),
+ status=payload.get("status"),
+ ).run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ except AIChatThreadInvalidError as ex:
+ return self.response_422(message=str(ex))
+ return self.response(200, result=_thread_dict(thread))
+
+ @expose("/thread/<thread_uuid>", methods=("DELETE",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.delete_thread",
+ log_to_statsd=False,
+ )
+ def delete_thread(self, thread_uuid: str) -> Response:
+ """Delete a conversation and its messages.
+ ---
+ delete:
+ summary: Delete a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ responses:
+ 200:
+ description: Conversation deleted
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import DeleteAIChatThreadCommand
+
+ try:
+ DeleteAIChatThreadCommand(thread_uuid, self._user_id()).run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ return self.response(200, message="OK")
+
+ @expose("/thread/<thread_uuid>/message", methods=("POST",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.post_message",
+ log_to_statsd=False,
+ )
+ def post_message(self, thread_uuid: str) -> Response:
+ """Post a user message and start a run.
+ ---
+ post:
+ summary: Post a message
+ description: >
+ Stores the user's message, creates a placeholder assistant message,
+ and starts a run. Returns immediately; consume the answer from the
+ stream endpoint using the returned run identifier.
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/MessagePostSchema'
+ responses:
+ 202:
+ description: Run accepted
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/RunAcceptedResponseSchema'
+ 400:
+ $ref: '#/components/responses/400'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ 422:
+ $ref: '#/components/responses/422'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.orchestrator import new_run_id
+ from superset.commands.ai import AppendAIChatMessageCommand
+
+ try:
+ payload = MessagePostSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ user_id = self._user_id()
+
+ try:
+ user_message = AppendAIChatMessageCommand(
+ thread_uuid,
+ user_id,
+ MessageRole.USER,
+ payload["content"],
+ request_id=payload.get("request_id"),
+ ).run()
+ # Created up front so a client that reconnects before any token
+ # arrives still has a row to attach its stream to.
+ assistant_message_command = AppendAIChatMessageCommand(
+ thread_uuid,
+ user_id,
+ MessageRole.ASSISTANT,
+ "",
+ request_id=payload.get("request_id"),
+ status=MessageStatus.PENDING,
+ )
+ assistant_message = assistant_message_command.run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ except (AIChatMessageInvalidError, AIChatThreadInvalidError) as ex:
+ return self.response_422(message=str(ex))
+
+ if assistant_message_command.created:
+ run_id = new_run_id()
+ _record_run_context(assistant_message, run_id, payload)
+ self._start_run(
+ thread_uuid=thread_uuid,
+ user_id=user_id,
+ run_id=run_id,
+ assistant_message_uuid=str(assistant_message.uuid),
+ agent_key=payload.get("agent_key"),
+ model=payload.get("model"),
+ page_context=payload.get("page_context"),
+ )
+ else:
+ # A retried idempotency key returns the run already attached to
+ # the reused assistant row instead of starting inference again.
+ run_id = str(
+ assistant_message.extra.get("run_id") or assistant_message.uuid
+ )
+
+ return self.response(
+ 202,
+ result={
+ "message_uuid": str(user_message.uuid),
+ "assistant_message_uuid": str(assistant_message.uuid),
+ "run_id": run_id,
+ },
+ )
+
+ @expose("/thread/<thread_uuid>/stream", methods=("GET",))
+ @protect()
+ @statsd_metrics
+ @permission_name("read")
+ def stream(self, thread_uuid: str) -> Response:
+ """Stream a run's events.
+ ---
+ get:
+ summary: Stream assistant events
+ description: >
+ Server-sent events for one run. Frame names are session, thinking,
+ thoughts, checkpoint, assistant_delta, final, error, cancelled and
+ done. The done frame is always last and reports whether the run
+ succeeded.
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ - in: query
+ name: run_id
+ required: true
+ schema:
+ type: string
+ responses:
+ 200:
+ description: An event stream
+ content:
+ text/event-stream:
+ schema:
+ type: string
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ # No @safe here: once headers are flushed an exception can no longer
+ # become a status code, so failures are reported as in-band error
frames.
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import AIChatMessageDAO, AIChatThreadDAO
+
+ run_id = request.args.get("run_id")
+ if not run_id:
+ return self.response_400(message="run_id is required")
+
+ # Ownership is checked before the stream opens; the run identifier
alone
+ # must not grant access to another user's conversation.
+ thread = AIChatThreadDAO.find_by_uuid_for_user(thread_uuid,
self._user_id())
+ if thread is None:
+ return self.response_404()
+
+ pending = _find_run_message(AIChatMessageDAO.find_for_thread(thread),
run_id)
+ if pending is None:
+ return self.response_404()
+
+ turn = None
+ if current_app.config.get("AI_ASSISTANT_EXECUTION_MODE") != "worker":
+ from superset.ai.orchestrator import TurnRequest
+
+ extra = pending.extra
+ turn = TurnRequest(
+ thread_uuid=thread_uuid,
+ user_id=self._user_id(),
+ run_id=run_id,
+ assistant_message_uuid=str(pending.uuid),
+ profile_key=extra.get("agent_key"),
+ model=extra.get("model"),
+ page_context=extra.get("page_context"),
+ )
+
+ generator = self._build_stream(run_id, turn)
+ response = Response(
+ generator,
+ content_type="text/event-stream; charset=utf-8",
+ headers={
+ "Cache-Control": "no-cache, no-transform",
+ "Connection": "keep-alive",
+ # Defeats proxy buffering, which otherwise holds frames until
+ # the response completes and makes streaming pointless.
+ "X-Accel-Buffering": "no",
+ "Content-Encoding": "identity",
+ },
+ direct_passthrough=False,
+ )
+ response.implicit_sequence_conversion = False
+ return response
+
+ @expose("/thread/<thread_uuid>/cancel", methods=("POST",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ def cancel(self, thread_uuid: str) -> Response:
+ """Ask a run to stop.
+ ---
+ post:
+ summary: Cancel a run
+ description: >
+ Cancellation is cooperative: the run stops at its next step
+ boundary. A run inside a single long model call or query will not
+ stop until that call returns.
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/CancelPostSchema'
+ responses:
+ 200:
+ description: Cancellation recorded
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.orchestrator import request_cancel
+ from superset.daos.ai import AIChatThreadDAO
+
+ try:
+ payload = CancelPostSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ if AIChatThreadDAO.find_by_uuid_for_user(thread_uuid, self._user_id())
is None:
+ return self.response_404()
+
+ request_cancel(payload["run_id"])
Review Comment:
This only verifies that the caller owns `thread_uuid`; it never verifies
that `run_id` belongs to that thread. A user can submit a run ID from another
conversation alongside their own thread and set that run's cancellation flag.
Should this resolve the run within the owned thread before calling
`request_cancel`?
##########
superset/ai/api.py:
##########
@@ -0,0 +1,996 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""
+REST API for the AI assistant.
+
+Every route carries ``@protect()`` and is reached through ``@expose`` on a
+``BaseSupersetApi`` subclass, which is what makes Flask-AppBuilder's
+authorization actually run. Ownership is enforced a second time in the command
+and DAO layers, so a conversation identifier is never on its own a capability.
+"""
+
+from __future__ import annotations
+
+import logging
+import time
+from collections.abc import Generator
+from typing import Any, cast
+
+from flask import current_app, request, Response, stream_with_context
+from flask_appbuilder.api import expose, permission_name, protect, safe
+from marshmallow import ValidationError
+
+from superset.ai.events import (
+ error_event,
+ KEEPALIVE_FRAME,
+ KEEPALIVE_INTERVAL_SECONDS,
+)
+from superset.ai.schemas import (
+ AgentResponseSchema,
+ CancelPostSchema,
+ FeedbackPostSchema,
+ MessagePostSchema,
+ RunAcceptedResponseSchema,
+ SuggestedPromptsPostSchema,
+ ThreadDetailResponseSchema,
+ ThreadPostSchema,
+ ThreadPutSchema,
+ ThreadResponseSchema,
+)
+from superset.ai.types import MessageRole, MessageStatus
+from superset.commands.ai.exceptions import (
+ AIChatMessageInvalidError,
+ AIChatMessageNotFoundError,
+ AIChatThreadInvalidError,
+ AIChatThreadNotFoundError,
+)
+from superset.extensions import event_logger
+from superset.utils.core import get_user_id
+from superset.utils.decorators import transaction
+from superset.views.base_api import BaseSupersetApi, statsd_metrics
+
+logger = logging.getLogger(__name__)
+
+#: Upper bound on how long a client may hold a stream open, so an abandoned
+#: browser tab cannot pin a worker indefinitely.
+_STREAM_TIMEOUT_SECONDS = 900
+
+#: How often a reader checks the event bus for new frames.
+#:
+#: Deliberately separate from ``KEEPALIVE_INTERVAL_SECONDS``. Passing the
+#: keep-alive interval as the poll interval made the reader sleep fifteen
seconds
+#: between checks and then deliver everything that had accumulated in one
batch —
+#: so a worker-mode run showed no streaming at all: the answer and every tool
call
+#: appeared in fifteen-second lumps. One controls responsiveness, the other how
+#: often an idle connection is reassured; they are not the same number.
+_EVENT_POLL_SECONDS = 0.1
+
+
+class AIRestApi(BaseSupersetApi):
+ """Conversations with the AI assistant."""
+
+ resource_name = "ai"
+ openapi_spec_tag = "AI Assistant"
+ allow_browser_login = True
+ class_permission_name = "AIAssistant"
+
+ openapi_spec_component_schemas = (
+ AgentResponseSchema,
+ CancelPostSchema,
+ FeedbackPostSchema,
+ MessagePostSchema,
+ RunAcceptedResponseSchema,
+ SuggestedPromptsPostSchema,
+ ThreadDetailResponseSchema,
+ ThreadPostSchema,
+ ThreadPutSchema,
+ ThreadResponseSchema,
+ )
+
+ @expose("/agent/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def agents(self) -> Response:
+ """List agent profiles the current user may select.
+ ---
+ get:
+ summary: List available agent profiles
+ responses:
+ 200:
+ description: Available profiles
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ type: array
+ items:
+ $ref: '#/components/schemas/AgentResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 403:
+ $ref: '#/components/responses/403'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.factories import get_profiles
+
+ profiles = get_profiles().visible_to_current_user()
+ return self.response(200, result=[p.to_public_dict() for p in
profiles])
+
+ @expose("/model/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def models(self) -> Response:
+ """List models this deployment has configured.
+ ---
+ get:
+ summary: List selectable models
+ responses:
+ 200:
+ description: Configured model identifiers
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ type: array
+ items:
+ type: string
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.factories import get_provider
+
+ return self.response(200, result=get_provider().available_models())
+
+ @expose("/thread/", methods=("POST",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.post_thread",
+ log_to_statsd=False,
+ )
+ def post_thread(self) -> Response:
+ """Create a conversation.
+ ---
+ post:
+ summary: Create a conversation
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/ThreadPostSchema'
+ responses:
+ 201:
+ description: Conversation created
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/ThreadResponseSchema'
+ 400:
+ $ref: '#/components/responses/400'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import CreateAIChatThreadCommand
+
+ try:
+ payload = ThreadPostSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ try:
+ thread = CreateAIChatThreadCommand(
+ user_id=self._user_id(),
+ title=payload.get("title"),
+ agent_key=payload.get("agent_key"),
+ ).run()
+ except AIChatThreadInvalidError as ex:
+ return self.response_422(message=str(ex))
+ return self.response(201, result=_thread_dict(thread))
+
+ @expose("/thread/", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def get_threads(self) -> Response:
+ """List the current user's conversations.
+ ---
+ get:
+ summary: List conversations
+ parameters:
+ - in: query
+ name: limit
+ schema:
+ type: integer
+ - in: query
+ name: offset
+ schema:
+ type: integer
+ responses:
+ 200:
+ description: Conversations
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ count:
+ type: integer
+ result:
+ type: array
+ items:
+ $ref: '#/components/schemas/ThreadResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import AIChatThreadDAO
+
+ limit = request.args.get("limit", type=int) or 50
+ offset = request.args.get("offset", type=int) or 0
+ threads = AIChatThreadDAO.find_all_for_user(
+ self._user_id(), limit=limit, offset=offset
+ )
+ return self.response(
+ 200,
+ count=len(threads),
+ result=[_thread_dict(thread) for thread in threads],
+ )
+
+ @expose("/thread/<thread_uuid>", methods=("GET",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("read")
+ def get_thread(self, thread_uuid: str) -> Response:
+ """Fetch a conversation and its messages.
+ ---
+ get:
+ summary: Get a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ responses:
+ 200:
+ description: Conversation with messages
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/ThreadDetailResponseSchema'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import (
+ AIChatFeedbackDAO,
+ AIChatMessageDAO,
+ AIChatThreadDAO,
+ )
+
+ user_id = self._user_id()
+ thread = AIChatThreadDAO.find_by_uuid_for_user(thread_uuid, user_id)
+ if thread is None:
+ return self.response_404()
+
+ messages = AIChatMessageDAO.find_for_thread(thread)
+ # Resolved for the whole transcript at once so the panel can show which
+ # replies this user already rated; without it a reload loses the
verdict
+ # and the message looks unrated.
+ verdicts = AIChatFeedbackDAO.find_verdicts_for_user(
+ [message.id for message in messages], user_id
+ )
+ detail = _thread_dict(thread)
+ detail["messages"] = [
+ _message_dict(message, liked=verdicts.get(message.id))
+ for message in messages
+ ]
+ return self.response(200, result=detail)
+
+ @expose("/thread/<thread_uuid>", methods=("PUT",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.put_thread",
+ log_to_statsd=False,
+ )
+ def put_thread(self, thread_uuid: str) -> Response:
+ """Rename or archive a conversation.
+ ---
+ put:
+ summary: Update a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/ThreadPutSchema'
+ responses:
+ 200:
+ description: Conversation updated
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ 422:
+ $ref: '#/components/responses/422'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import UpdateAIChatThreadCommand
+
+ try:
+ payload = ThreadPutSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ try:
+ thread = UpdateAIChatThreadCommand(
+ thread_uuid,
+ self._user_id(),
+ title=payload.get("title"),
+ status=payload.get("status"),
+ ).run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ except AIChatThreadInvalidError as ex:
+ return self.response_422(message=str(ex))
+ return self.response(200, result=_thread_dict(thread))
+
+ @expose("/thread/<thread_uuid>", methods=("DELETE",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.delete_thread",
+ log_to_statsd=False,
+ )
+ def delete_thread(self, thread_uuid: str) -> Response:
+ """Delete a conversation and its messages.
+ ---
+ delete:
+ summary: Delete a conversation
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ responses:
+ 200:
+ description: Conversation deleted
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.commands.ai import DeleteAIChatThreadCommand
+
+ try:
+ DeleteAIChatThreadCommand(thread_uuid, self._user_id()).run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ return self.response(200, message="OK")
+
+ @expose("/thread/<thread_uuid>/message", methods=("POST",))
+ @protect()
+ @safe
+ @statsd_metrics
+ @permission_name("write")
+ @event_logger.log_this_with_context(
+ action=lambda self, *args, **kwargs:
f"{self.__class__.__name__}.post_message",
+ log_to_statsd=False,
+ )
+ def post_message(self, thread_uuid: str) -> Response:
+ """Post a user message and start a run.
+ ---
+ post:
+ summary: Post a message
+ description: >
+ Stores the user's message, creates a placeholder assistant message,
+ and starts a run. Returns immediately; consume the answer from the
+ stream endpoint using the returned run identifier.
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ requestBody:
+ content:
+ application/json:
+ schema:
+ $ref: '#/components/schemas/MessagePostSchema'
+ responses:
+ 202:
+ description: Run accepted
+ content:
+ application/json:
+ schema:
+ type: object
+ properties:
+ result:
+ $ref: '#/components/schemas/RunAcceptedResponseSchema'
+ 400:
+ $ref: '#/components/responses/400'
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ 422:
+ $ref: '#/components/responses/422'
+ """
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.ai.orchestrator import new_run_id
+ from superset.commands.ai import AppendAIChatMessageCommand
+
+ try:
+ payload = MessagePostSchema().load(request.json or {})
+ except ValidationError as error:
+ return self.response_400(message=error.messages)
+ user_id = self._user_id()
+
+ try:
+ user_message = AppendAIChatMessageCommand(
+ thread_uuid,
+ user_id,
+ MessageRole.USER,
+ payload["content"],
+ request_id=payload.get("request_id"),
+ ).run()
+ # Created up front so a client that reconnects before any token
+ # arrives still has a row to attach its stream to.
+ assistant_message_command = AppendAIChatMessageCommand(
+ thread_uuid,
+ user_id,
+ MessageRole.ASSISTANT,
+ "",
+ request_id=payload.get("request_id"),
+ status=MessageStatus.PENDING,
+ )
+ assistant_message = assistant_message_command.run()
+ except AIChatThreadNotFoundError:
+ return self.response_404()
+ except (AIChatMessageInvalidError, AIChatThreadInvalidError) as ex:
+ return self.response_422(message=str(ex))
+
+ if assistant_message_command.created:
+ run_id = new_run_id()
+ _record_run_context(assistant_message, run_id, payload)
+ self._start_run(
+ thread_uuid=thread_uuid,
+ user_id=user_id,
+ run_id=run_id,
+ assistant_message_uuid=str(assistant_message.uuid),
+ agent_key=payload.get("agent_key"),
+ model=payload.get("model"),
+ page_context=payload.get("page_context"),
+ )
+ else:
+ # A retried idempotency key returns the run already attached to
+ # the reused assistant row instead of starting inference again.
+ run_id = str(
+ assistant_message.extra.get("run_id") or assistant_message.uuid
+ )
+
+ return self.response(
+ 202,
+ result={
+ "message_uuid": str(user_message.uuid),
+ "assistant_message_uuid": str(assistant_message.uuid),
+ "run_id": run_id,
+ },
+ )
+
+ @expose("/thread/<thread_uuid>/stream", methods=("GET",))
+ @protect()
+ @statsd_metrics
+ @permission_name("read")
+ def stream(self, thread_uuid: str) -> Response:
+ """Stream a run's events.
+ ---
+ get:
+ summary: Stream assistant events
+ description: >
+ Server-sent events for one run. Frame names are session, thinking,
+ thoughts, checkpoint, assistant_delta, final, error, cancelled and
+ done. The done frame is always last and reports whether the run
+ succeeded.
+ parameters:
+ - in: path
+ name: thread_uuid
+ required: true
+ schema:
+ type: string
+ format: uuid
+ - in: query
+ name: run_id
+ required: true
+ schema:
+ type: string
+ responses:
+ 200:
+ description: An event stream
+ content:
+ text/event-stream:
+ schema:
+ type: string
+ 401:
+ $ref: '#/components/responses/401'
+ 404:
+ $ref: '#/components/responses/404'
+ """
+ # No @safe here: once headers are flushed an exception can no longer
+ # become a status code, so failures are reported as in-band error
frames.
+ if (unavailable := self._reject_if_unconfigured()) is not None:
+ return unavailable
+
+ from superset.daos.ai import AIChatMessageDAO, AIChatThreadDAO
+
+ run_id = request.args.get("run_id")
+ if not run_id:
+ return self.response_400(message="run_id is required")
+
+ # Ownership is checked before the stream opens; the run identifier
alone
+ # must not grant access to another user's conversation.
+ thread = AIChatThreadDAO.find_by_uuid_for_user(thread_uuid,
self._user_id())
+ if thread is None:
+ return self.response_404()
+
+ pending = _find_run_message(AIChatMessageDAO.find_for_thread(thread),
run_id)
+ if pending is None:
+ return self.response_404()
+
+ turn = None
+ if current_app.config.get("AI_ASSISTANT_EXECUTION_MODE") != "worker":
Review Comment:
A reconnect or duplicate GET for an already completed inline run constructs
a fresh `TurnRequest` and runs inference/tools again, so the same run ID can
double-charge or repeat an external tool effect. Could this atomically claim
only pending work and replay persisted events for terminal runs?
##########
superset/ai/orchestrator.py:
##########
@@ -0,0 +1,647 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+"""
+Runs one assistant turn end to end.
+
+Sits between the HTTP layer and the runtime: loads the conversation, assembles
+the prompt, resolves the tools the chosen profile allows, drives the runtime,
+publishes every event to the bus, and records the outcome on the assistant
+message.
+
+Deliberately independent of *where* it runs. The same function body serves the
+inline path and the Celery path, which is what makes the execution mode a
+configuration choice rather than two implementations that drift apart.
+"""
+
+from __future__ import annotations
+
+import asyncio
+import logging
+import uuid as uuid_module
+from collections.abc import AsyncIterator, Iterator
+from dataclasses import dataclass
+from typing import Any
+
+from superset.ai.events import (
+ cancelled_event,
+ done_event,
+ error_event,
+ GENERIC_ERROR_MESSAGE,
+ session_event,
+ StreamEvent,
+)
+from superset.ai.llm.base import Message
+from superset.ai.telemetry import bind_run, current_run, start_run
+from superset.ai.types import MessageRole, MessageStatus, RunOutcome,
StreamEventType
+from superset.utils.decorators import transaction
+
+logger = logging.getLogger(__name__)
+
+#: Cache key prefix for a run's cancellation flag. A flag rather than a signal
+#: because a worker cannot be interrupted mid-call reliably; the runtime checks
+#: this between steps.
+_CANCEL_PREFIX = "ai-cancel-"
+
+#: How long a cancellation request stays meaningful.
+_CANCEL_TTL_SECONDS = 900
+
+#: Stored when a run is stopped before it produced any answer, so the
+#: transcript still records that the turn happened.
+_STOPPED_WITHOUT_ANSWER = "_Stopped before an answer was produced._"
+
+#: Stored when a run exhausted its time budget without saying anything. Phrased
+#: as something the user can act on, because retrying is usually the right
move.
+_TIMED_OUT_WITHOUT_ANSWER = (
+ "The assistant ran out of time before it could answer. Please try again."
+)
+
+#: Ceiling on the page context recorded on a message. Well below the prompt's
own
+#: limit: this is stored per turn and read back with the whole transcript.
+_RECORDED_CONTEXT_LIMIT = 4_000
+
+
+@dataclass
+class TurnRequest:
+ """One unit of work: answer the latest message on a thread."""
+
+ thread_uuid: str
+ user_id: int
+ run_id: str
+ #: Assistant message row to fill in. Created before the run starts so a
+ #: client that reconnects has something to attach to.
+ assistant_message_uuid: str
+ profile_key: str | None = None
+ #: Concrete model to pin, overriding the profile's tier.
+ model: str | None = None
+ #: What the user had on screen when they asked. Supplied by the client,
+ #: which is the only party that knows which tab is open, what is typed in
+ #: the editor and which filters are applied.
+ page_context: dict[str, Any] | None = None
+
+ def to_payload(self) -> dict[str, Any]:
+ """Serialise for the task broker."""
+ return {
+ "thread_uuid": self.thread_uuid,
+ "user_id": self.user_id,
+ "run_id": self.run_id,
+ "assistant_message_uuid": self.assistant_message_uuid,
+ "profile_key": self.profile_key,
+ "model": self.model,
+ "page_context": self.page_context,
+ }
+
+ @classmethod
+ def from_payload(cls, payload: dict[str, Any]) -> TurnRequest:
+ """Rebuild from a broker payload."""
+ return cls(**payload)
+
+
+def new_run_id() -> str:
+ """Identifier for one run, used as the event-stream key."""
+ return str(uuid_module.uuid4())
+
+
+#: Runs cancelled in this process.
+#:
+#: Held alongside the cache rather than instead of it. Superset's default cache
+#: is a null cache, which accepts a write and discards it — so a cache-only
+#: implementation would leave cancellation silently broken on a default
install,
+#: with the button appearing to work and nothing stopping. This set makes
inline
+#: execution correct with no cache at all; the cache is what carries a
+#: cancellation across processes for worker execution.
+_CANCELLED_LOCALLY: set[str] = set()
+
+
+def request_cancel(run_id: str) -> None:
+ """
+ Ask a run to stop.
+
+ Cooperative by design: the flag is recorded here and observed by the
runtime
+ between steps. A run blocked inside a single long model call or query will
+ not notice until that call returns, which is a real limit worth documenting
+ rather than hiding.
+ """
+ from superset.extensions import cache_manager
+
+ _CANCELLED_LOCALLY.add(run_id)
+ try:
+ cache_manager.cache.set(
+ f"{_CANCEL_PREFIX}{run_id}", True, timeout=_CANCEL_TTL_SECONDS
+ )
+ except Exception: # pylint: disable=broad-except
+ logger.warning("Could not record cancellation for AI run %s", run_id)
+
+
+def is_cancelled(run_id: str) -> bool:
+ """Whether a stop has been requested for this run."""
+ from superset.extensions import cache_manager
+
+ if run_id in _CANCELLED_LOCALLY:
+ return True
+ try:
+ return bool(cache_manager.cache.get(f"{_CANCEL_PREFIX}{run_id}"))
+ except Exception: # pylint: disable=broad-except
+ # A cache that cannot be read must not make every run appear cancelled;
+ # that would stop all inference the moment the cache went away.
+ return False
+
+
+def clear_cancel(run_id: str) -> None:
+ """Drop a run's cancellation flag."""
+ from superset.extensions import cache_manager
+
+ _CANCELLED_LOCALLY.discard(run_id)
+ try:
+ cache_manager.cache.delete(f"{_CANCEL_PREFIX}{run_id}")
+ except Exception: # pylint: disable=broad-except
+ logger.debug("Could not clear cancellation flag for AI run %s", run_id)
+
+
+def stream_turn(request: TurnRequest) -> Iterator[StreamEvent]:
+ """
+ Answer a turn, yielding events as they happen.
+
+ This is the primary entry point. Inline execution consumes it directly from
+ inside the streaming response, which means the producer and the reader are
+ the same process by construction — important because Superset runs several
+ web workers, and a turn that published to one process's in-memory queue
+ while the browser's stream landed on another would appear to hang forever.
+
+ Never raises for an operational failure: a failure is an ``error`` event
and
+ an ``error`` message status, because the caller may already have flushed
+ response headers or may be a worker with no one to report to.
+ """
+ recorder = start_run(
+ run_id=request.run_id,
+ thread_uuid=request.thread_uuid,
+ user_id=request.user_id,
+ )
+ # Shared with ``_run`` so the ``finally`` below can see the runtime's
+ # partial result and whether the message was already written.
+ state: dict[str, Any] = {}
+ try:
+ # Bound here rather than inside ``_run`` so that a run which fails
before
+ # it has resolved a profile still produces a start and an end, and so
+ # that the runtime can report its own spans without the runtime
contract
+ # growing a telemetry parameter.
+ with bind_run(recorder):
+ recorder.run_started()
+ yield from _run(request, state)
+ except Exception as ex: # pylint: disable=broad-except
+ logger.exception("AI turn failed for run %s", request.run_id)
+ recorder.error(ex)
+ recorder.run_ended(outcome=RunOutcome.ERROR)
+ answer, extra = _partial_from_state(state)
+ extra["outcome"] = RunOutcome.ERROR.value
+ _finalise_message(
+ request.assistant_message_uuid,
+ # The generic text rather than the exception: this is persisted and
+ # served back to the browser, so it must not carry internals. The
+ # detail is in the log line above, keyed by run id.
+ content=answer or GENERIC_ERROR_MESSAGE,
+ status=MessageStatus.ERROR,
+ extra=extra,
+ )
+ state["finalised"] = True
+ yield error_event()
+ yield done_event(ok=False)
+ finally:
+ clear_cancel(request.run_id)
+ # A client that stops the run, or simply navigates away, abandons this
+ # generator part-way through. Nothing above will have written the
+ # message, so it would otherwise sit in ``streaming`` with no content
+ # for ever — the user loses both the partial answer and any record that
+ # the turn happened. Persist whatever was produced.
+ _abandon_message(request.assistant_message_uuid, state)
+ # Idempotent, so the ordinary paths above win.
+ recorder.run_ended(outcome=RunOutcome.CANCELLED)
+
+
+def execute_turn(request: TurnRequest) -> RunOutcome:
+ """
+ Answer a turn, publishing events to the event bus.
+
+ Used by worker execution, where the reader is in another process. Shares
its
+ whole body with :func:`stream_turn` so the two execution modes cannot drift
+ apart in behaviour.
+ """
+ from superset.ai.eventbus import get_event_bus
+
+ bus = get_event_bus()
+ outcome = RunOutcome.SUCCESS
+
+ for event in stream_turn(request):
+ bus.publish(request.run_id, event)
+ if event.type is StreamEventType.ERROR:
+ outcome = RunOutcome.ERROR
+ elif event.type is StreamEventType.CANCELLED:
+ outcome = RunOutcome.CANCELLED
+ elif event.type is StreamEventType.DONE and not
event.payload.get("ok"):
+ # A run that ended un-ok without an explicit error frame timed out.
+ if outcome is RunOutcome.SUCCESS:
+ outcome = RunOutcome.TIMEOUT
+
+ return outcome
+
+
+def _run(request: TurnRequest, state: dict[str, Any]) -> Iterator[StreamEvent]:
+ """Assemble and drive the run. See :func:`stream_turn` for error policy."""
+ from superset.ai.factories import (
+ get_profiles,
+ get_provider,
+ get_runtime,
+ get_tools_for_profile,
+ )
+ from superset.ai.policy import load_policy_chain
+ from superset.ai.runtime.base import RunRequest
+ from superset.daos.ai import AIChatMessageDAO, AIChatThreadDAO
+
+ recorder = current_run()
+
+ thread = AIChatThreadDAO.find_by_uuid_for_user(request.thread_uuid,
request.user_id)
+ if thread is None:
+ # The thread vanished between accepting the message and running it.
+ recorder.run_ended(outcome=RunOutcome.ERROR)
+ yield error_event("That conversation is no longer available.")
+ yield done_event(ok=False)
+ return
+
+ profile = get_profiles().get(request.profile_key)
+ tools = get_tools_for_profile(profile)
+ provider = get_provider()
+ runtime = get_runtime(provider)
+ state["runtime"] = runtime
+
+ yield session_event(request.thread_uuid, request.assistant_message_uuid)
+
+ from superset.ai.page_context import render_page_context
+
+ history = _build_history(AIChatMessageDAO.find_for_thread(thread))
+ # Recorded as well as prompted with, so the transcript can show what the
+ # assistant was told about the user's screen. An answer that looks wrong is
+ # usually an answer to a different question than the reader assumed, and
the
+ # page context is where that difference lives.
+ rendered_context = render_page_context(request.page_context)
+ state["page_context"] = rendered_context
+ system_prompt = _build_system_prompt(tools, rendered_context)
+ model = _resolved_model(provider, request.model, profile)
+
+ recorder.describe(
+ agent_key=profile.key,
+ model=model,
+ question=_latest_question(history),
+ )
+
+ run_request = RunRequest(
+ messages=history,
+ system_prompt=system_prompt,
+ tools=tools,
+ policies=load_policy_chain(),
+ model_alias=profile.model_alias,
Review Comment:
The resolved `request.model` is recorded above, but only
`profile.model_alias` reaches `RunRequest`; the runtime therefore selects the
alias while persistence and telemetry identify the requested model. Should the
resolved pin be passed through to the completion request?
##########
superset-frontend/src/features/ai/hooks/useChatBot.ts:
##########
@@ -0,0 +1,1323 @@
+/**
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied. See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+/**
+ * @fileoverview Conversation state and the send loop.
+ *
+ * Runs are tracked per conversation, not globally. That is the point of the
+ * structure: a user can start something slow in one conversation, switch to
+ * another and keep working, and come back to find the first still going. A
single
+ * `isLoading` flag would have made switching away cancel or corrupt the run.
+ *
+ * The server owns the transcript. A finished run is re-read from it rather
than
+ * assembled from the frames, so the tool calls persisted on the message are
what
+ * the user sees, and what they see survives a reload.
+ */
+
+import { useCallback, useEffect, useRef, useState } from 'react';
+import type { TextAreaRef } from 'antd/es/input/TextArea';
+import { logging } from '@apache-superset/core/utils';
+import { t } from '@apache-superset/core/translation';
+import {
+ type AiAgent,
+ type AiToolCall,
+ type ChatMessageWithMeta,
+ type ChatTab,
+ type CheckpointPayload,
+} from '../types';
+import {
+ AGENT_STORAGE_KEY,
+ ChatRequestAbortedError,
+ ChatStreamEventError,
+ ChatStreamTimeoutError,
+ DEFAULT_AGENT_KEY,
+ DEFAULT_CHAT_AGENT,
+ cancelChatRun,
+ describeRequestError,
+ fetchAgents,
+ fetchSuggestedPrompts,
+ loadStoredAgentKey,
+ normalizeChatAgents,
+ startRun,
+ streamRun,
+ submitFeedback,
+} from './chatRequest';
+import {
+ NEW_CHAT_NAME,
+ createThread,
+ deleteThread as deleteThreadApi,
+ getThread,
+ listThreads,
+ threadToTab,
+ updateThread,
+} from './chatThreadsApi';
+import { buildQuickPrompts } from './quickPrompts';
+import {
+ buildPageContextPayload,
+ usePageContext,
+ type PageContext,
+} from './usePageContext';
+
+/** Cache of the conversation list, so the menu renders before the list
arrives. */
+export const CHAT_TABS_STORAGE_KEY = 'superset-chat-tabs';
+
+/** Which conversation was last open. */
+export const ACTIVE_TAB_STORAGE_KEY = 'superset-chat-active-tab';
+
+/** Recent inputs, recalled with the arrow keys. */
+export const HISTORY_STORAGE_KEY = 'superset-chat-history';
+
+export { AGENT_STORAGE_KEY } from './chatRequest';
+
+/** How many inputs the arrow-key history keeps. */
+const MAX_INPUT_HISTORY = 50;
+
+/** A conversation title derived from a message is clipped to this. */
+const MAX_TAB_NAME_LENGTH = 30;
+
+export type ChatRunStatus = 'running' | 'cancelling';
+
+/** Shared empty list, so a render with no steps yet keeps a stable identity.
*/
+const EMPTY_TOOL_CALLS: AiToolCall[] = [];
+
+interface ActiveChatRun {
+ requestId: string;
+ tabId: string;
+ threadId: string;
+ runId?: string;
+ controller: AbortController;
+ isStreaming: boolean;
+ liveThoughts: string;
+ liveToolLog: string;
+ /**
+ * Steps taken so far, as structured records rather than log lines.
+ *
+ * Carried alongside `liveToolLog` so a run in flight can be rendered the
same
+ * way a finished one is — expandable per step, with the SQL and the rows it
+ * returned — instead of as a wall of text that only becomes legible once the
+ * transcript is re-read from the server.
+ */
+ liveToolCalls: AiToolCall[];
+ /** The page context this run was given, so the live view can show it too. */
+ livePageContext?: string;
+ /**
+ * The answer so far, as the model produces it.
+ *
+ * Rendered directly: the deltas used to be folded into `liveThinking`, which
+ * nothing displayed, so an answer appeared in one piece the moment the run
+ * ended however long it had taken to generate.
+ */
+ liveAnswer: string;
+ liveThinking: string;
+ status: ChatRunStatus;
+ startedAt: number;
+ checkpoint: CheckpointPayload | null;
+}
+
+/**
+ * An identifier for a turn.
+ *
+ * Drawn from `crypto`, not `Math.random`. These become the idempotency key on
a
+ * turn and the handle used to cancel one, so a value another session could
guess
+ * is a correctness and a security problem rather than merely a collision risk.
+ */
+const generateId = (): string => {
+ if (typeof crypto.randomUUID === 'function') {
+ return crypto.randomUUID();
+ }
+ // Older engines expose the entropy source without the convenience wrapper.
+ const bytes = new Uint8Array(16);
+ crypto.getRandomValues(bytes);
+ return Array.from(bytes, byte => byte.toString(16).padStart(2,
'0')).join('');
+};
+
+const createNewTab = (name: string = NEW_CHAT_NAME): ChatTab => ({
+ id: generateId(),
+ name,
+ messages: [],
+ createdAt: Date.now(),
+});
+
+const truncateTabName = (
+ name: string,
+ maxLength: number = MAX_TAB_NAME_LENGTH,
+): string =>
+ name.length <= maxLength ? name : `${name.substring(0, maxLength)}...`;
+
+const readJson = <T>(key: string, fallback: T): T => {
+ try {
+ const stored = localStorage.getItem(key);
+ return stored ? (JSON.parse(stored) as T) : fallback;
+ } catch (caught) {
+ logging.warn(`[ai] could not read ${key}`, caught);
+ return fallback;
+ }
+};
+
+const writeJson = (key: string, value: unknown): void => {
+ try {
+ localStorage.setItem(key, JSON.stringify(value));
+ } catch (caught) {
+ logging.warn(`[ai] could not write ${key}`, caught);
+ }
+};
+
+/**
+ * Reconciles the server's transcript with what is already on screen.
+ *
+ * The server's copy is authoritative — it carries the tool calls — but it is
not
+ * necessarily complete the moment a run ends, and replacing outright would
then
+ * erase an answer the user has just read. So anything local that the server
has
+ * not accounted for is kept, matched by identity first and by role and content
+ * second, which is how a locally-appended turn is recognised once the server
+ * returns its own copy of it under a real uuid.
+ */
+export const mergeMessages = (
+ fromServer: ChatMessageWithMeta[],
+ local: ChatMessageWithMeta[],
+): ChatMessageWithMeta[] => {
+ const serverIds = new Set(fromServer.map(message => message.id));
+ const serverTurns = new Set(
+ fromServer.map(message => `${message.role}:${message.content}`),
+ );
+ const unaccounted = local.filter(
+ message =>
+ !serverIds.has(message.id) &&
+ !serverTurns.has(`${message.role}:${message.content}`),
+ );
+ return [...fromServer, ...unaccounted];
+};
+
+/**
+ * The `page_context` body for one turn.
+ *
+ * Returns undefined when there is nothing to send, so an omitted field is
+ * distinguishable from an empty one.
+ */
+export const buildRequestPageContext = (
+ context: PageContext | undefined,
+ directive?: string,
+): Record<string, unknown> | undefined => {
+ const payload = context ? buildPageContextPayload(context) : undefined;
+ if (!directive) {
+ return payload;
+ }
+ const existing = payload?.helper_directives;
+ return {
+ ...payload,
+ helper_directives: [
Review Comment:
`systemPrompt` is sent as `page_context.helper_directives`, but the server
renderer only reads the structured SQL/chart/dashboard/markdown fields. Calls
such as `triggerAIAction({ systemPrompt: ... })` therefore silently lose their
directive. Could the backend consume this field through a dedicated, trusted
prompt channel?
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]