sadpandajoe commented on code in PR #43133: URL: https://github.com/apache/superset/pull/43133#discussion_r3995801185
########## 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: [ + directive, + ...(Array.isArray(existing) ? existing : []), + ], + }; +}; + +export interface UseChatBotReturn { + // Conversations + chatTabs: ChatTab[]; + activeTabId: string; + activeTab: ChatTab | undefined; + threadsLoaded: boolean; + handleNewChat: () => Promise<string>; + handleSelectTab: (tabId: string) => Promise<void>; + handleDeleteTab: (tabId: string) => Promise<void>; + handleRenameTab: (tabId: string, newName: string) => void; + // Messages of the active conversation + messages: ChatMessageWithMeta[]; + // Input + inputValue: string; + setInputValue: (value: string) => void; + handleKeyDown: (event: React.KeyboardEvent) => void; + inputRef: React.RefObject<TextAreaRef>; + messagesEndRef: React.RefObject<HTMLDivElement>; + // The run in flight, if any, for the active conversation + isLoading: boolean; + isStreamingResponse: boolean; + liveThoughts: string; + liveToolLog: string; + /** Steps taken so far in the run in flight, for the structured live view. */ + liveToolCalls: AiToolCall[]; + /** The page context the run in flight was given. */ + livePageContext?: string; + /** The answer so far for the run in flight. */ + liveAnswer: string; + checkpoint: CheckpointPayload | null; + activeRunStatus: ChatRunStatus | null; + error?: string; + // Actions + sendMessage: ( + messageOverride?: string, + systemPromptOverride?: string, + ) => Promise<void>; + handleCancelRun: () => Promise<void>; + handleCheckpointContinue: () => void; + handleFeedback: (messageId: string, feedback: 'like' | 'dislike') => void; + messageFeedback: Record<string, 'like' | 'dislike'>; + // Suggestions + /** The message whose run just ended; its thought process stays open. */ + justCompletedId?: string; + quickPrompts: string[]; + loadQuickPrompts: () => void; + applyQuickPrompt: (prompt: string) => Promise<void>; + // Agent profiles + agents: AiAgent[]; + selectedAgent: string; + setSelectedAgent: (key: string) => void; + // Page context + pageContext: PageContext; + includePageContext: boolean; + toggleIncludePageContext: () => void; +} + +export const useChatBot = (): UseChatBotReturn => { + const [chatTabs, setChatTabs] = useState<ChatTab[]>(() => + readJson<ChatTab[]>(CHAT_TABS_STORAGE_KEY, []).map(tab => ({ + // The cache is a placeholder for the menu; message bodies are re-read from + // the server so a stale cache cannot show a conversation that has moved on. + ...tab, + messages: [], + })), + ); + const [activeTabId, setActiveTabId] = useState<string>(() => { + try { + return localStorage.getItem(ACTIVE_TAB_STORAGE_KEY) ?? ''; + } catch { + return ''; + } + }); + const [threadsLoaded, setThreadsLoaded] = useState(false); + const [error, setError] = useState<string | undefined>(undefined); + + const [inputValue, setInputValue] = useState(''); + const [activeRunsByTab, setActiveRunsByTab] = useState< + Record<string, ActiveChatRun> + >({}); + const [quickPrompts, setQuickPrompts] = useState<string[]>([]); + const [messageFeedback, setMessageFeedback] = useState< + Record<string, 'like' | 'dislike'> + >({}); + const [includePageContext, setIncludePageContext] = useState(true); + /** + * The assistant message whose run has only just ended. + * + * Its thought process stays open, because collapsing it the instant the answer + * lands moves everything below it — the answer the user is mid-sentence through + * jumps up the panel. Older messages start closed. + */ + const [justCompletedId, setJustCompletedId] = useState<string | undefined>(); + const [agents, setAgents] = useState<AiAgent[]>([DEFAULT_CHAT_AGENT]); + const [selectedAgent, setSelectedAgent] = useState<string>(() => + loadStoredAgentKey(AGENT_STORAGE_KEY), + ); + + const [messageHistory, setMessageHistory] = useState<string[]>(() => + readJson<string[]>(HISTORY_STORAGE_KEY, []), + ); + const [historyIndex, setHistoryIndex] = useState(-1); + const [currentDraft, setCurrentDraft] = useState(''); + + const messagesEndRef = useRef<HTMLDivElement>(null); + const inputRef = useRef<TextAreaRef>(null); + + /** + * The run map and the conversation list are also held in refs, and the refs are + * the authority. + * + * The send loop has to ask "is this still my run?" between awaits, and it cannot + * ask React: a run that starts and fails inside one batch never causes a render, + * so a ref synced at render time would still be empty and the loop would discard + * its own result as stale. Writing the ref at the point of mutation removes that + * window. The callbacks read the refs rather than the state so their identities + * do not churn on every streamed frame, which would restart effects mid-run. + */ + const activeRunsByTabRef = useRef<Record<string, ActiveChatRun>>({}); + const chatTabsRef = useRef<ChatTab[]>(chatTabs); + const activeTabIdRef = useRef(activeTabId); + activeTabIdRef.current = activeTabId; + + const updateRuns = useCallback( + ( + updater: ( + previous: Record<string, ActiveChatRun>, + ) => Record<string, ActiveChatRun>, + ) => { + activeRunsByTabRef.current = updater(activeRunsByTabRef.current); + setActiveRunsByTab(activeRunsByTabRef.current); + }, + [], + ); + + const updateTabs = useCallback( + (updater: (previous: ChatTab[]) => ChatTab[]) => { + chatTabsRef.current = updater(chatTabsRef.current); + setChatTabs(chatTabsRef.current); + }, + [], + ); + + /** Resolved when the user answers a checkpoint; see `streamRun`. */ + const checkpointGateRef = useRef<{ resolve: () => void } | null>(null); Review Comment: The checkpoint gate is still global while run controllers are per tab, so a confirmation or cancel action from one tab can release the gate created by another tab. Could checkpoint gates be keyed by thread or run like the controllers? ########## superset/ai/api.py: ########## @@ -0,0 +1,989 @@ +# 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 = AppendAIChatMessageCommand( + thread_uuid, + user_id, + MessageRole.ASSISTANT, + "", + request_id=payload.get("request_id"), + status=MessageStatus.PENDING, + ).run() + except AIChatThreadNotFoundError: + return self.response_404() + except (AIChatMessageInvalidError, AIChatThreadInvalidError) as ex: + return self.response_422(message=str(ex)) + + run_id = new_run_id() Review Comment: `request_id` still suppresses only duplicate message insertion; every replay allocates and schedules a new run. The user message, assistant placeholder, run context, and queue submission are also separate commits/effects, so a mid-sequence or ambiguous broker failure leaves a partial turn and a retry can duplicate work. Could this use one durable, claimable run and return that same run for an existing request ID? ########## superset/ai/eventbus.py: ########## @@ -0,0 +1,314 @@ +# 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. +""" +Carries streamed events from whatever produced them to the HTTP response. + +Two implementations, matching the two execution modes. Inline execution needs +nothing more than an in-process queue. Worker execution needs a shared, +*replayable* channel — replayable because a browser that loses its connection +must be able to rejoin a run already in progress, which rules out +publish/subscribe: a subscriber that was absent when an event was published +never sees it. + +The Redis implementation therefore uses streams, and reuses the cache backend +that Superset's async-query channel already configures rather than introducing +a second Redis client to operate. +""" + +from __future__ import annotations + +import logging +import queue +from abc import ABC, abstractmethod +from collections.abc import Iterator +from typing import Any + +from superset.ai.events import StreamEvent +from superset.ai.types import StreamEventType +from superset.utils import json + +logger = logging.getLogger(__name__) + +#: Yielded by :meth:`BaseEventBus.consume` when nothing arrived within the poll +#: interval, so a caller can emit a keep-alive rather than block indefinitely. +IDLE = None + +#: Terminal event types. Seeing one ends consumption, so a reader does not hang +#: waiting for a producer that has already finished. +_TERMINAL = frozenset( + {StreamEventType.DONE, StreamEventType.ERROR, StreamEventType.CANCELLED} +) + + +class BaseEventBus(ABC): + """A per-run channel of events.""" + + @abstractmethod + def publish(self, run_id: str, event: StreamEvent) -> None: + """Append an event to a run's channel.""" + + @abstractmethod + def consume( + self, + run_id: str, + timeout_seconds: float, + poll_seconds: float = 1.0, + ) -> Iterator[StreamEvent | None]: + """ + Yield a run's events until a terminal one arrives or time runs out. + + Yields :data:`IDLE` when a poll interval passes with nothing new, which + is the caller's cue to send a keep-alive frame. + """ + + @abstractmethod + def close(self, run_id: str) -> None: + """Release any resources held for a run.""" + + +class MemoryEventBus(BaseEventBus): + """ + An in-process queue per run. + + Correct only when the producer and the streaming request share a process. + Selecting this alongside worker execution would leave every stream silent, + which :func:`get_event_bus` refuses to allow. + """ + + def __init__(self) -> None: + self._queues: dict[str, queue.SimpleQueue[StreamEvent]] = {} + + def _queue_for(self, run_id: str) -> queue.SimpleQueue[StreamEvent]: + return self._queues.setdefault(run_id, queue.SimpleQueue()) + + def publish(self, run_id: str, event: StreamEvent) -> None: + self._queue_for(run_id).put(event) + + def consume( + self, + run_id: str, + timeout_seconds: float, + poll_seconds: float = 1.0, + ) -> Iterator[StreamEvent | None]: + import time + + # Deliberately not ``_queue_for``: reading must not create a channel. + # This bus lives for the life of the process, so a client polling + # unknown run identifiers would otherwise grow the dict without bound. + channel = self._queues.get(run_id) + deadline = time.monotonic() + timeout_seconds + + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + return + if channel is None: + # The producer may not have published yet; look again rather + # than deciding the run does not exist. Only report idle if it + # is still absent, so a channel that appeared during the wait + # is drained on this pass instead of costing an extra tick. + channel = self._queues.get(run_id) + if channel is None: + yield IDLE + time.sleep(min(poll_seconds, remaining)) + continue + try: + # Bounded by whichever is sooner, so a generous poll interval + # cannot overshoot the caller's deadline. + event = channel.get(timeout=min(poll_seconds, remaining)) + except queue.Empty: + yield IDLE + continue + yield event + if event.type in _TERMINAL: + return + + def close(self, run_id: str) -> None: + self._queues.pop(run_id, None) + + +class RedisStreamEventBus(BaseEventBus): + """ + A Redis stream per run. + + Replayable by construction: a reconnecting reader starts from the beginning + of the stream and catches up, which is what makes worker execution usable + from a browser on a flaky connection. + """ + + def __init__( + self, + cache: Any, + prefix: str = "ai-events-", + ttl_seconds: int = 900, + ) -> None: + self._cache = cache + self._prefix = prefix + self._ttl = ttl_seconds + + def _stream(self, run_id: str) -> str: + return f"{self._prefix}{run_id}" + + def publish(self, run_id: str, event: StreamEvent) -> None: + payload = { + "data": json.dumps({"type": event.type.value, "payload": event.payload}) + } + # A failure to publish must not kill the run that is producing useful + # work; the reader will time out and the answer is still persisted. + try: + self._cache.xadd(self._stream(run_id), payload, "*", 10_000) Review Comment: Publishing still only appends events; the expiry is set from reader cleanup, so a disconnected client can leave an unread stream key without any TTL. Could the key receive or refresh its expiry when events are published? ########## superset/ai/api.py: ########## @@ -0,0 +1,989 @@ +# 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 = AppendAIChatMessageCommand( + thread_uuid, + user_id, + MessageRole.ASSISTANT, + "", + request_id=payload.get("request_id"), + status=MessageStatus.PENDING, + ).run() + except AIChatThreadNotFoundError: + return self.response_404() + except (AIChatMessageInvalidError, AIChatThreadInvalidError) as ex: + return self.response_422(message=str(ex)) + + 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"), + ) + + 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 in inline mode still constructs and executes a fresh `TurnRequest` on every GET instead of attaching to the existing run. Could reconnects reuse durable run state so an SSE retry cannot execute the turn twice? ########## superset/ai/runtime/messages.py: ########## @@ -0,0 +1,574 @@ +# 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. +""" +The default runtime: a plain tool-use loop over the provider's message API. + +Chosen as the default because it needs nothing beyond an HTTP call — no agent +engine subprocess, no working directory, no bundled binary — so it works with +whatever provider a deployment configures. +""" + +from __future__ import annotations + +import logging +import time +from collections.abc import AsyncIterator +from typing import Any + +from superset.ai.events import ( + assistant_delta_event, + checkpoint_event, + error_event, + final_event, + GENERIC_ERROR_MESSAGE, + StreamEvent, + thinking_event, + thoughts_event, +) +from superset.ai.llm.base import ( + CompletionRequest, + LLMError, + LLMResponse, + Message, + StreamEventKind, + ToolCall, + ToolResult, +) +from superset.ai.runtime.base import BaseAgentRuntime, RunRequest, RunResult +from superset.ai.telemetry import ( + current_run, + POLICY_DENIED, + RunRecorder, + TOOL_UNAVAILABLE, +) +from superset.ai.types import MessageRole, ProgressStage, TokenUsage + +logger = logging.getLogger(__name__) + +#: How much of a tool's output is kept on the persisted message. The model +#: still sees the whole thing; this is the audit copy. +_RECORDED_OUTPUT_LIMIT = 2_000 + +#: Size of the chunks the finished answer is delivered in. +_DELIVERY_CHUNK_SIZE = 512 + +#: How much reasoning is kept on the result. Reasoning can run several times +#: longer than the answer, and this is persisted next to it. +_RECORDED_THOUGHTS_LIMIT = 8_000 + +_NO_ANSWER = ( + "I wasn't able to reach an answer for that. Try narrowing the question, " + "or naming the dataset you have in mind." +) + + +class MessagesApiRuntime(BaseAgentRuntime): + """ + Alternates model calls and tool calls until the model stops asking. + + Two behaviours are worth understanding before changing this class. + + First, prose the model emits *before* a tool call is treated as reasoning, + not answer: it becomes a ``thoughts`` event and is dropped from the answer. + A model narrating "the orders table looks right, let me check" is stating a + hypothesis it may abandon, and appending that to the answer produces a + reply that contradicts itself. + + Second, the loop always terminates and never raises for an operational + failure. By the time it runs, response headers have been flushed and an + exception can no longer become an HTTP status, so every failure is an event. + """ + + def __init__(self, provider: Any) -> None: + super().__init__(provider) + self._result = RunResult() + #: Set when the model signals it has finished answering. + self._finished = False + #: The most recent round trip's response, or ``None`` if it failed. The + #: turn methods are generators and cannot return a value. + self._last_response: LLMResponse | None = None + #: Whether any answer text has already been sent as it was generated. The + #: finished answer is only replayed in chunks when it has not. + self._streamed_text = False + + @property + def result(self) -> RunResult: + return self._result + + async def run(self, request: RunRequest) -> AsyncIterator[StreamEvent]: + self._result = RunResult() + self._finished = False + self._last_response = None + self._streamed_text = False + answer_parts: list[str] = [] + + yield thinking_event(ProgressStage.START, "Working on your question") + + # The provider's connection pool belongs to the loop this run is driven + # on, and the caller closes that loop as soon as the run ends. Closing + # here — inside the loop, however the run finishes, including when the + # generator is abandoned mid-way by a user pressing stop — is what keeps + # a client from being finalised against a dead loop. + try: + async for event in self._turn_loop(request, answer_parts): + yield event + + # A run that failed or was abandoned has already said so; emitting an + # answer as well would contradict it. + if self._result.error is not None or self._result.cancelled: + return + + answer = "\n\n".join(part for part in answer_parts if part).strip() + self._result.answer = answer or _NO_ANSWER + + # Only replayed when nothing was streamed — a provider without + # streaming support still gets to deliver its answer progressively. + # Replaying after live text would show the answer twice. + if not self._streamed_text: + for chunk in _chunk(self._result.answer): + yield assistant_delta_event(chunk) + yield final_event(self._result.answer) + finally: + await self.provider.aclose() + + async def _turn_loop( + self, + request: RunRequest, + answer_parts: list[str], + ) -> AsyncIterator[StreamEvent]: + """ + Alternate model and tool calls until the model stops or a budget runs out. + + Appends to ``answer_parts`` rather than returning the answer, because an + async generator cannot both yield events and return a value. + """ + deadline = time.monotonic() + request.timeout_seconds + conversation = list(request.messages) + + for turn in range(1, request.max_turns + 1): + self._result.turns = turn + + if self._should_stop(request, deadline): + if self._result.timed_out: + yield thinking_event( + ProgressStage.FALLBACK, + "Taking longer than expected — answering with what I have", + ) + return + + async for event in self._safe_turn(request, conversation, turn): + yield event + response = self._last_response + if response is None: + yield error_event() + return + + async for event in self._consume( + request, response, conversation, answer_parts + ): + yield event + + if self._finished or self._result.cancelled: + return + + # Budget exhausted without the model choosing to stop. + yield thinking_event( + ProgressStage.FALLBACK, + "Reached the step limit — answering with what I have", + ) + + async def _consume( + self, + request: RunRequest, + response: LLMResponse, + conversation: list[Message], + answer_parts: list[str], + ) -> AsyncIterator[StreamEvent]: + """Act on one model response, running any tools it asked for.""" + if response.thinking: + self._record_thoughts(response.thinking) + yield thoughts_event(response.thinking) + + if not response.wants_tools: + self._finished = True + if response.text: + answer_parts.append(response.text) + # Recorded as it arrives, not just at the end, so a run stopped + # after this point still persists what the user already saw. + self._result.answer = "\n\n".join( + part for part in answer_parts if part + ).strip() + return + + # Prose accompanying a tool call is reasoning, not answer. + if response.text: + self._record_thoughts(response.text) + yield thoughts_event(response.text) + + conversation.append( + Message( + role=MessageRole.ASSISTANT, + content=response.text, + tool_calls=list(response.tool_calls), + ) + ) + + results: list[ToolResult] = [] + async for event in self._run_tools(request, response.tool_calls, results): + yield event + + conversation.append(Message(role=MessageRole.USER, tool_results=results)) + + async def _run_tools( + self, + request: RunRequest, + calls: list[ToolCall], + results: list[ToolResult], + ) -> AsyncIterator[StreamEvent]: + """Execute this turn's tool calls, appending outcomes to ``results``.""" + for call in calls: + if self._cancelled(request): + self._result.cancelled = True + return + + yield thinking_event( + ProgressStage.TOOL, + f"Running {call.name}", + {"tool_name": call.name}, + ) + result, detail = self._invoke_tool(request, call) + results.append(result) + record = self._record_call(call, result, detail) + + # The frame carries the same record that is persisted, rather than a + # subset assembled separately. The subset was missing the arguments + # and the output, so a step expanded during a run showed nothing at + # all unless its tool happened to supply a display — and then filled + # itself in on reload, which looked like the detail arrived late. + # Sharing one record makes that class of drift impossible. + yield checkpoint_event( + f"{'Failed' if result.is_error else 'Finished'} {call.name}", + # ``tool_name`` as well as ``name``: the progress frames use that + # key, so a consumer reading either finds what it expects. + {"tool_name": call.name, **record}, + ) + + async def _safe_turn( + self, + request: RunRequest, + conversation: list[Message], + turn: int, + ) -> AsyncIterator[StreamEvent]: + """ + One model round trip, converting failure into a ``None`` response. + + A generator rather than a coroutine so the answer can reach the client as + the model produces it. The response is handed back on + :attr:`_last_response` because an async generator cannot both yield events + and return a value — the same reason ``_turn_loop`` writes into + ``answer_parts``. + + The failure detail goes to the log; the caller emits a message that cannot + leak a URL, a credential or a fragment of someone else's query. + """ + recorder = current_run() + started = time.monotonic() + self._last_response = None + try: + async for event in self._one_turn(request, conversation): + yield event + except LLMError as ex: + logger.warning("AI provider error on turn %s: %s", turn, ex) + self._result.error = str(ex) + self._trace_model_call(recorder, request, turn, started, error=ex) + self._last_response = None + return + except Exception as ex: # pylint: disable=broad-except + logger.exception("Unexpected error in AI runtime on turn %s", turn) + self._result.error = GENERIC_ERROR_MESSAGE + self._trace_model_call(recorder, request, turn, started, error=ex) + self._last_response = None + return + self._trace_model_call( + recorder, request, turn, started, response=self._last_response + ) + + def _trace_model_call( + self, + recorder: RunRecorder, + request: RunRequest, + turn: int, + started: float, + response: LLMResponse | None = None, + error: BaseException | None = None, + ) -> None: + """ + Report one round trip to telemetry. + + Content is passed as-is; whether any of it survives into a trace is the + redaction policy's decision, made in one place rather than here. + """ + if not recorder.enabled: + return + usage = response.usage if response is not None else TokenUsage() + recorder.model_call( + turn=turn, + # The concrete identifier when the provider reported one, and the + # capability tier otherwise, so a trace can always be grouped by + # what the run asked for. + model=usage.get("model") or request.model_alias.value, + duration_ms=int((time.monotonic() - started) * 1000), + input_tokens=usage.get("input_tokens"), + output_tokens=usage.get("output_tokens"), + stop_reason=response.stop_reason if response is not None else None, + error_type=type(error).__name__ if error is not None else None, + system_prompt=request.system_prompt, + response_text=response.text if response is not None else None, + ) + if error is not None: + recorder.error(error) + + async def _one_turn( + self, + request: RunRequest, + conversation: list[Message], + ) -> AsyncIterator[StreamEvent]: + """ + Call the model once, yielding answer text as the model produces it. + + Streaming is used when the provider supports it. The assembled response + is left on :attr:`_last_response` rather than returned, because a + generator cannot do both; it has the same shape either way, so callers do + not branch on which path ran. + """ + completion = CompletionRequest( Review Comment: The exact configured `model` is still omitted from `CompletionRequest`; only `model_alias` reaches the provider. When an operator pins a model behind an alias, the runtime can select the provider default instead. Could the exact model be propagated? ########## superset/ai/runtime/messages.py: ########## @@ -0,0 +1,574 @@ +# 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. +""" +The default runtime: a plain tool-use loop over the provider's message API. + +Chosen as the default because it needs nothing beyond an HTTP call — no agent +engine subprocess, no working directory, no bundled binary — so it works with +whatever provider a deployment configures. +""" + +from __future__ import annotations + +import logging +import time +from collections.abc import AsyncIterator +from typing import Any + +from superset.ai.events import ( + assistant_delta_event, + checkpoint_event, + error_event, + final_event, + GENERIC_ERROR_MESSAGE, + StreamEvent, + thinking_event, + thoughts_event, +) +from superset.ai.llm.base import ( + CompletionRequest, + LLMError, + LLMResponse, + Message, + StreamEventKind, + ToolCall, + ToolResult, +) +from superset.ai.runtime.base import BaseAgentRuntime, RunRequest, RunResult +from superset.ai.telemetry import ( + current_run, + POLICY_DENIED, + RunRecorder, + TOOL_UNAVAILABLE, +) +from superset.ai.types import MessageRole, ProgressStage, TokenUsage + +logger = logging.getLogger(__name__) + +#: How much of a tool's output is kept on the persisted message. The model +#: still sees the whole thing; this is the audit copy. +_RECORDED_OUTPUT_LIMIT = 2_000 + +#: Size of the chunks the finished answer is delivered in. +_DELIVERY_CHUNK_SIZE = 512 + +#: How much reasoning is kept on the result. Reasoning can run several times +#: longer than the answer, and this is persisted next to it. +_RECORDED_THOUGHTS_LIMIT = 8_000 + +_NO_ANSWER = ( + "I wasn't able to reach an answer for that. Try narrowing the question, " + "or naming the dataset you have in mind." +) + + +class MessagesApiRuntime(BaseAgentRuntime): + """ + Alternates model calls and tool calls until the model stops asking. + + Two behaviours are worth understanding before changing this class. + + First, prose the model emits *before* a tool call is treated as reasoning, + not answer: it becomes a ``thoughts`` event and is dropped from the answer. + A model narrating "the orders table looks right, let me check" is stating a + hypothesis it may abandon, and appending that to the answer produces a + reply that contradicts itself. + + Second, the loop always terminates and never raises for an operational + failure. By the time it runs, response headers have been flushed and an + exception can no longer become an HTTP status, so every failure is an event. + """ + + def __init__(self, provider: Any) -> None: + super().__init__(provider) + self._result = RunResult() + #: Set when the model signals it has finished answering. + self._finished = False + #: The most recent round trip's response, or ``None`` if it failed. The + #: turn methods are generators and cannot return a value. + self._last_response: LLMResponse | None = None + #: Whether any answer text has already been sent as it was generated. The + #: finished answer is only replayed in chunks when it has not. + self._streamed_text = False + + @property + def result(self) -> RunResult: + return self._result + + async def run(self, request: RunRequest) -> AsyncIterator[StreamEvent]: + self._result = RunResult() + self._finished = False + self._last_response = None + self._streamed_text = False + answer_parts: list[str] = [] + + yield thinking_event(ProgressStage.START, "Working on your question") + + # The provider's connection pool belongs to the loop this run is driven + # on, and the caller closes that loop as soon as the run ends. Closing + # here — inside the loop, however the run finishes, including when the + # generator is abandoned mid-way by a user pressing stop — is what keeps + # a client from being finalised against a dead loop. + try: + async for event in self._turn_loop(request, answer_parts): + yield event + + # A run that failed or was abandoned has already said so; emitting an + # answer as well would contradict it. + if self._result.error is not None or self._result.cancelled: + return + + answer = "\n\n".join(part for part in answer_parts if part).strip() + self._result.answer = answer or _NO_ANSWER + + # Only replayed when nothing was streamed — a provider without + # streaming support still gets to deliver its answer progressively. + # Replaying after live text would show the answer twice. + if not self._streamed_text: + for chunk in _chunk(self._result.answer): + yield assistant_delta_event(chunk) + yield final_event(self._result.answer) + finally: + await self.provider.aclose() + + async def _turn_loop( + self, + request: RunRequest, + answer_parts: list[str], + ) -> AsyncIterator[StreamEvent]: + """ + Alternate model and tool calls until the model stops or a budget runs out. + + Appends to ``answer_parts`` rather than returning the answer, because an + async generator cannot both yield events and return a value. + """ + deadline = time.monotonic() + request.timeout_seconds + conversation = list(request.messages) + + for turn in range(1, request.max_turns + 1): + self._result.turns = turn + + if self._should_stop(request, deadline): + if self._result.timed_out: + yield thinking_event( + ProgressStage.FALLBACK, + "Taking longer than expected — answering with what I have", + ) + return + + async for event in self._safe_turn(request, conversation, turn): + yield event + response = self._last_response + if response is None: + yield error_event() + return + + async for event in self._consume( + request, response, conversation, answer_parts + ): + yield event + + if self._finished or self._result.cancelled: + return + + # Budget exhausted without the model choosing to stop. + yield thinking_event( + ProgressStage.FALLBACK, + "Reached the step limit — answering with what I have", + ) + + async def _consume( + self, + request: RunRequest, + response: LLMResponse, + conversation: list[Message], + answer_parts: list[str], + ) -> AsyncIterator[StreamEvent]: + """Act on one model response, running any tools it asked for.""" + if response.thinking: + self._record_thoughts(response.thinking) + yield thoughts_event(response.thinking) + + if not response.wants_tools: + self._finished = True + if response.text: + answer_parts.append(response.text) + # Recorded as it arrives, not just at the end, so a run stopped + # after this point still persists what the user already saw. + self._result.answer = "\n\n".join( + part for part in answer_parts if part + ).strip() + return + + # Prose accompanying a tool call is reasoning, not answer. + if response.text: + self._record_thoughts(response.text) + yield thoughts_event(response.text) + + conversation.append( + Message( + role=MessageRole.ASSISTANT, + content=response.text, + tool_calls=list(response.tool_calls), + ) + ) + + results: list[ToolResult] = [] + async for event in self._run_tools(request, response.tool_calls, results): + yield event + + conversation.append(Message(role=MessageRole.USER, tool_results=results)) + + async def _run_tools( + self, + request: RunRequest, + calls: list[ToolCall], + results: list[ToolResult], + ) -> AsyncIterator[StreamEvent]: + """Execute this turn's tool calls, appending outcomes to ``results``.""" + for call in calls: + if self._cancelled(request): + self._result.cancelled = True + return + + yield thinking_event( + ProgressStage.TOOL, + f"Running {call.name}", + {"tool_name": call.name}, + ) + result, detail = self._invoke_tool(request, call) + results.append(result) + record = self._record_call(call, result, detail) + + # The frame carries the same record that is persisted, rather than a + # subset assembled separately. The subset was missing the arguments + # and the output, so a step expanded during a run showed nothing at + # all unless its tool happened to supply a display — and then filled + # itself in on reload, which looked like the detail arrived late. + # Sharing one record makes that class of drift impossible. + yield checkpoint_event( + f"{'Failed' if result.is_error else 'Finished'} {call.name}", + # ``tool_name`` as well as ``name``: the progress frames use that + # key, so a consumer reading either finds what it expects. + {"tool_name": call.name, **record}, + ) + + async def _safe_turn( + self, + request: RunRequest, + conversation: list[Message], + turn: int, + ) -> AsyncIterator[StreamEvent]: + """ + One model round trip, converting failure into a ``None`` response. + + A generator rather than a coroutine so the answer can reach the client as + the model produces it. The response is handed back on + :attr:`_last_response` because an async generator cannot both yield events + and return a value — the same reason ``_turn_loop`` writes into + ``answer_parts``. + + The failure detail goes to the log; the caller emits a message that cannot + leak a URL, a credential or a fragment of someone else's query. + """ + recorder = current_run() + started = time.monotonic() + self._last_response = None + try: + async for event in self._one_turn(request, conversation): + yield event + except LLMError as ex: + logger.warning("AI provider error on turn %s: %s", turn, ex) + self._result.error = str(ex) + self._trace_model_call(recorder, request, turn, started, error=ex) + self._last_response = None + return + except Exception as ex: # pylint: disable=broad-except + logger.exception("Unexpected error in AI runtime on turn %s", turn) + self._result.error = GENERIC_ERROR_MESSAGE + self._trace_model_call(recorder, request, turn, started, error=ex) + self._last_response = None + return + self._trace_model_call( + recorder, request, turn, started, response=self._last_response + ) + + def _trace_model_call( + self, + recorder: RunRecorder, + request: RunRequest, + turn: int, + started: float, + response: LLMResponse | None = None, + error: BaseException | None = None, + ) -> None: + """ + Report one round trip to telemetry. + + Content is passed as-is; whether any of it survives into a trace is the + redaction policy's decision, made in one place rather than here. + """ + if not recorder.enabled: + return + usage = response.usage if response is not None else TokenUsage() + recorder.model_call( + turn=turn, + # The concrete identifier when the provider reported one, and the + # capability tier otherwise, so a trace can always be grouped by + # what the run asked for. + model=usage.get("model") or request.model_alias.value, + duration_ms=int((time.monotonic() - started) * 1000), + input_tokens=usage.get("input_tokens"), + output_tokens=usage.get("output_tokens"), + stop_reason=response.stop_reason if response is not None else None, + error_type=type(error).__name__ if error is not None else None, + system_prompt=request.system_prompt, + response_text=response.text if response is not None else None, + ) + if error is not None: + recorder.error(error) + + async def _one_turn( + self, + request: RunRequest, + conversation: list[Message], + ) -> AsyncIterator[StreamEvent]: + """ + Call the model once, yielding answer text as the model produces it. + + Streaming is used when the provider supports it. The assembled response + is left on :attr:`_last_response` rather than returned, because a + generator cannot do both; it has the same shape either way, so callers do + not branch on which path ran. + """ + completion = CompletionRequest( + messages=conversation, + system=request.system_prompt, + model_alias=request.model_alias, + tools=tuple(request.tools.definitions()) if request.tools else (), + ) + + if not self.provider.supports_streaming: Review Comment: `runtime/messages.py` still calls `provider.complete` and `provider.stream` directly, so normal chat completions bypass the configured retry policy. The deadline is also checked only before the call, so one slow or hung provider request can exceed the turn timeout indefinitely. Could these calls go through a wrapper that enforces both retry policy and the remaining deadline? ########## superset/ai/tools/authoring.py: ########## @@ -0,0 +1,310 @@ +# 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. +"""Native AI adapters for Superset's existing MCP authoring tools.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Callable +from importlib import import_module +from threading import Thread +from typing import Any, ClassVar, TypeVar + +from pydantic import BaseModel, ValidationError + +from superset.ai.tools.base import AITool, ToolError, ToolOutput +from superset.mcp_service.chart.schemas import GenerateChartRequest +from superset.mcp_service.dashboard.schemas import GenerateDashboardRequest +from superset.mcp_service.dataset.schemas import CreateVirtualDatasetRequest +from superset.utils import json + +ModelT = TypeVar("ModelT", bound=BaseModel) +ToolCaller = Callable[[BaseModel], Any] + +_MCP_TOOL_MODULES = { + "create_virtual_dataset": ( + "superset.mcp_service.dataset.tool.create_virtual_dataset" + ), + "generate_chart": "superset.mcp_service.chart.tool.generate_chart", + "generate_dashboard": ("superset.mcp_service.dashboard.tool.generate_dashboard"), +} + + +def _tool_schema(model: type[BaseModel]) -> dict[str, Any]: + """Expose a request model without its server-only warning field.""" + schema = model.model_json_schema() + properties = dict(schema.get("properties", {})) + properties.pop("sanitization_warnings", None) + schema["properties"] = properties + if required := schema.get("required"): + schema["required"] = [ + name for name in required if name != "sanitization_warnings" + ] + return schema + + +def _validate(model: type[ModelT], payload: dict[str, Any], label: str) -> ModelT: + """Turn Pydantic errors into a correction the model can act on.""" + try: + return model.model_validate(payload) + except ValidationError as ex: + issues = [] + for error in ex.errors(include_url=False)[:3]: + location = ".".join(str(part) for part in error["loc"]) + issues.append(f"{location}: {error['msg']}") + raise ToolError(f"Invalid {label} request: {'; '.join(issues)}.") from ex + + +def _payload(response: Any) -> dict[str, Any]: + if isinstance(response, BaseModel): + return response.model_dump(mode="json", exclude_none=True) + if isinstance(response, dict): + return response + raise ToolError("Superset returned an unexpected authoring response.") + + +async def _call_mcp_tool(tool_name: str, request: BaseModel) -> Any: + """Call the registered tool through FastMCP so it gets a real context.""" + import_module(_MCP_TOOL_MODULES[tool_name]) + + from fastmcp import Client + + from superset.mcp_service.app import mcp + + arguments = { + "request": request.model_dump( + mode="json", + exclude={"sanitization_warnings"}, + exclude_none=True, + ) + } + async with Client(mcp) as client: + result = await client.call_tool(tool_name, arguments) + + if result.is_error: + raise ToolError(f"Superset could not run {tool_name}.") + return ( + result.structured_content + if result.structured_content is not None + else result.data + ) + + +def _run_mcp_tool(tool_name: str, request: BaseModel) -> dict[str, Any]: + """Run FastMCP off the agent loop with isolated Flask request state.""" + from flask import current_app, g + + try: + app = current_app._get_current_object() + user = getattr(g, "user", None) + except RuntimeError as ex: + raise ToolError("Authoring requires an authenticated request.") from ex + + username = getattr(user, "username", None) + email = getattr(user, "email", None) + if not username and not email: + raise ToolError("Authoring requires an authenticated user.") + + outcome: dict[str, Any] = {} + + def run() -> None: + try: + from flask import g as worker_g + + from superset.mcp_service.auth import load_user_with_relationships + + with app.test_request_context(): + worker_g.user = load_user_with_relationships( + username=str(username) if username else None, + email=str(email) if email else None, + ) + if worker_g.user is None: + raise ToolError("The authenticated user could not be reloaded.") + outcome["value"] = asyncio.run(_call_mcp_tool(tool_name, request)) + except BaseException as ex: # noqa: BLE001 + outcome["error"] = ex + + worker = Thread(target=run, name="superset-ai-authoring", daemon=True) + worker.start() + worker.join(float(app.config.get("AI_AGENT_TIMEOUT_SECONDS", 300))) + + if worker.is_alive(): + raise ToolError("Superset authoring timed out.") Review Comment: The call still returns an error while the mutating thread is live, so the model can retry before the first asset becomes visible and create a duplicate. It also always waits the global timeout, so a shorter profile deadline cannot be enforced. Could this receive the remaining run deadline and contain or durably reconcile the ambiguous outcome before another authoring call is allowed? -- 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]
