diff --git a/pkg-py/src/commons/_agent.py b/pkg-py/src/commons/_agent.py index d852107a..468eb274 100644 --- a/pkg-py/src/commons/_agent.py +++ b/pkg-py/src/commons/_agent.py @@ -11,6 +11,7 @@ import copy import warnings +import weakref from collections.abc import AsyncGenerator, Mapping, Sequence from typing import Any, Literal, NoReturn @@ -29,9 +30,12 @@ from ._context_layer import ContextLayer, augment_context_layer from ._data_source import DataSource from ._definitions import Registry, build_registry +from ._execution._driver import Worker +from ._execution._thread import WorkerThread from ._handles import HandleStore from ._measures import SemanticLayer, resolve_injections, semantic_layer from ._prompt import ( + EXECUTION_TOOL, check_instructions, read_instructions, render_system_prompt, @@ -40,6 +44,7 @@ ) from ._provenance import collect_appended_tags, derive_provenance_tag, provenance_aside from ._reminders import append_restored_conversation_reminder, append_turn_reminder +from ._run_python import run_python_description, run_python_tools from ._tools import FirstTouch, ToolContext, build_commons_tools __all__ = ["Commons"] @@ -84,11 +89,22 @@ class Commons(Chat[Any, Any]): `## Additional instructions` heading at the end of commons' built-in system prompt, as a string or the path to a text or Markdown file. + The agent can run Python code that its model writes. The code runs in a + separate, sandboxed Python process, which starts the first time the model + runs code. `network` sets whether that process can reach the network: + `"none"` (the default) or `"full"`. The sandbox works on Linux and macOS. + For local development on another system, set the + `COMMONS_ALLOW_UNSAFE_FALLBACK` environment variable to run the code with + limited checks instead. Warning! These checks are not intended to be a + security boundary. + Construction raises a TypeError if `client` is not a `chatlas.Chat`, if an entry of `data_sources` is not a `DataSource`, or if a layer is not the layer its argument claims; a ValueError if `data_sources` names no - source or a measure asks for an injection no named source can fill; and - a FileNotFoundError if `instructions` names a file that does not exist. + source, a measure asks for an injection no named source can fill, or + `network` is neither `"none"` nor `"full"`; a FileNotFoundError if + `instructions` names a file that does not exist; and a RuntimeError if + this host cannot sandbox the session and the opt-in is not set. """ def __init__( @@ -99,6 +115,7 @@ def __init__( context_layer: ContextLayer | None = None, *, instructions: str | None = None, + network: Literal["none", "full"] = "none", ) -> None: if not isinstance(client, Chat): raise TypeError( @@ -124,8 +141,14 @@ def __init__( f"{type(semantic_layer).__name__}." ) check_instructions(instructions) + # Built here so a host that cannot sandbox the session fails now, + # before any model asks to run code. The process starts on first use. + worker = Worker( + network=network, + measure_sources=list(semantic_layer.source_text.values()), + ) - # Share the provider, which carries the chosen model; shallow-copy + # Share the provider, which knows the chosen model; shallow-copy # the chat kwargs so later changes don't cross between the two. super().__init__( provider=client.provider, kwargs_chat=copy.copy(client.kwargs_chat) @@ -152,18 +175,33 @@ def __init__( ) self._restore_reminder_pending = False - tools = build_commons_tools( - ToolContext( - sources=sources, - measures=self._measures, - definitions=self._definitions, - context_layer=self._context_layer, - handles=self._handles, - citation_request=self._citation_request, - injections=self._injections, - first_touch=self._first_touch, - ) + context = ToolContext( + sources=sources, + measures=self._measures, + definitions=self._definitions, + context_layer=self._context_layer, + handles=self._handles, + citation_request=self._citation_request, + injections=self._injections, + first_touch=self._first_touch, + ) + tools = build_commons_tools(context) + self._python = WorkerThread(worker) + # Close the session's process and thread when the agent is + # collected or the interpreter exits. The callback names only the + # WorkerThread, so the finalizer never keeps the agent alive. + weakref.finalize(self, self._python.close) + self._run_python, self._run_python_sync = run_python_tools( + self._python, + context, + run_python_description( + [tool.name for tool in tools], + has_measures=bool(self._measures), + network=network, + ), + network, ) + tools.append(self._run_python) self.set_tools(list(tools)) self.system_prompt = _system_prompt( sources, @@ -204,7 +242,19 @@ def chat( was_pending = self._restore_reminder_pending inputs = self._prepare_turn_inputs(args) self._citation_request.reset() - response = super().chat(*inputs, echo=echo, stream=stream, kwargs=kwargs) + # A synchronous chat cannot await an async tool, so chatlas + # refuses to run while the async run_python tool is registered. + # This workaround swaps in the sync variant of run_python for this call. + swap = any(tool.name == EXECUTION_TOOL for tool in self.get_tools()) + if swap: + self.register_tool(self._run_python_sync, force=True) + try: + response = super().chat( + *inputs, echo=echo, stream=stream, kwargs=kwargs + ) + finally: + if swap: + self.register_tool(self._run_python, force=True) self._consume_restore_reminder(was_pending) return response diff --git a/pkg-py/src/commons/_display.py b/pkg-py/src/commons/_display.py index d36a2706..220342e8 100644 --- a/pkg-py/src/commons/_display.py +++ b/pkg-py/src/commons/_display.py @@ -24,6 +24,7 @@ _URL = re.compile(r"https?://", re.IGNORECASE) __all__ = [ + "CODE_ANALYSIS", "CONTEXT_SEARCH", "DATA_RETRIEVAL", "DISPLAY_EXTRA_KEY", @@ -60,6 +61,7 @@ class Title: CONTEXT_SEARCH = Title("Searching context", "Searched context") TABLE_INSPECTION = Title("Inspecting a table", "Inspected a table") DATA_RETRIEVAL = Title("Retrieving data", "Retrieved data") +CODE_ANALYSIS = Title("Analyzing data", "Analyzed data") def tool_display( diff --git a/pkg-py/src/commons/_execution/_driver.py b/pkg-py/src/commons/_execution/_driver.py index 67e3d253..8a3df010 100644 --- a/pkg-py/src/commons/_execution/_driver.py +++ b/pkg-py/src/commons/_execution/_driver.py @@ -98,6 +98,9 @@ def __init__( self._last_used = 0.0 self._lock = asyncio.Lock() self._ids = itertools.count(1) + # Set when a cancelled call shut the session down; take_restart() + # reports it once, so the model learns its variables were reset. + self._restarted = False async def __aenter__(self) -> Self: return self @@ -136,6 +139,15 @@ async def run( self._pending -= 1 self._schedule_reap() + def take_restart(self) -> bool: + """Whether a cancelled call restarted the session since the last ask. + + Only a call cancelled while it ran counts; one cancelled while it + waited for the lock left the session as it was. + """ + restarted, self._restarted = self._restarted, False + return restarted + async def aclose(self) -> None: """Close the worker, cancelling the idle reap. @@ -203,6 +215,7 @@ async def _call( # still running the code, or about to answer into a channel the # next call would misread as its own reply. Shut it down before # releasing the lock. + self._restarted = True await self._shutdown() raise except Exception as exc: # noqa: BLE001 - any start or write failure fails the call diff --git a/pkg-py/src/commons/_execution/_thread.py b/pkg-py/src/commons/_execution/_thread.py new file mode 100644 index 00000000..64b54911 --- /dev/null +++ b/pkg-py/src/commons/_execution/_thread.py @@ -0,0 +1,159 @@ +"""Runs the session's ``Worker`` in a background thread with its own event loop. + +chatlas runs an async tool only from ``stream_async()``, and runs a sync tool +on the caller's event loop, where a long call would stop every other task on +that loop. With the ``Worker`` in its own thread, an async caller can wait for +a call without blocking its loop, and a sync caller can block on the same +session. Variables therefore persist whether the agent is used through +``chat()`` or ``stream_async()``. +""" + +from __future__ import annotations + +import asyncio +import concurrent.futures +import threading + +from .._handles import HandleStore +from ._driver import Failure, Worker +from ._protocol import Error, Result + +__all__ = ["WorkerThread"] + +# How long close() waits for the worker's shutdown, which is itself bounded +# by its grace periods once the calls in flight are cancelled. +CLOSE_TIMEOUT = 30.0 + +_CLOSED = Failure(message="the Python session is closed.") + +_Reply = Result | Error | Failure + + +class WorkerThread: + """Runs ``worker`` in a background thread, which starts on the first call.""" + + def __init__(self, worker: Worker) -> None: + self._worker = worker + self._loop: asyncio.AbstractEventLoop | None = None + self._thread: threading.Thread | None = None + # Guards the lifecycle: a call is either submitted before close() + # begins, and so cancelled by it, or refused as closed. + self._lock = threading.Lock() + self._closed = False + self._calls: set[concurrent.futures.Future[_Reply]] = set() + + async def run(self, code: str, handles: HandleStore | None = None) -> _Reply: + """Run ``code`` without blocking the caller's event loop. + + If the caller is cancelled, the call is cancelled too, and the worker + shuts down as it does when ``Worker.run`` is cancelled. + """ + future = self._submit(code, handles) + if future is None: + return _CLOSED + try: + return await asyncio.wrap_future(future) + except asyncio.CancelledError: + task = asyncio.current_task() + if self._closed and future.cancelled() and not (task and task.cancelling()): + return _CLOSED + raise + + def run_sync(self, code: str, handles: HandleStore | None = None) -> _Reply: + """Run ``code``, blocking the calling thread until the reply arrives.""" + if threading.current_thread() is self._thread: + raise RuntimeError("run_sync() would deadlock on the worker's own loop") + future = self._submit(code, handles) + if future is None: + return _CLOSED + try: + return future.result() + except concurrent.futures.CancelledError: + if self._closed: + return _CLOSED + raise + except BaseException: + # Ctrl-C while waiting stops the call, as cancelling an async + # caller does, rather than leaving it to run out its timeout. + future.cancel() + raise + + def take_restart(self) -> bool: + """Whether a cancelled call restarted the session since the last ask.""" + return self._worker.take_restart() + + def close(self) -> None: + """Close the worker, then stop the loop and its thread. + + Calling it again does nothing. Running calls are cancelled first, so + the worker stops quickly instead of waiting for a call to time out. If + the shutdown takes longer than ``CLOSE_TIMEOUT``, it continues in the + background, and the loop stops when it ends. When called from the + loop's own thread, it starts the shutdown and returns at once. + """ + with self._lock: + if self._closed: + return + self._closed = True + loop, thread = self._loop, self._thread + calls = list(self._calls) + if loop is None or thread is None: + return + for call in calls: + call.cancel() + closing = asyncio.run_coroutine_threadsafe(self._worker.aclose(), loop) + closing.add_done_callback(lambda _: loop.call_soon_threadsafe(loop.stop)) + # Waiting here on the loop's own thread would block the shutdown. + if threading.current_thread() is thread: + return + try: + closing.result(CLOSE_TIMEOUT) + except TimeoutError: + return + finally: + if closing.done(): + thread.join(CLOSE_TIMEOUT) + + def _submit( + self, code: str, handles: HandleStore | None + ) -> concurrent.futures.Future[_Reply] | None: + """Queue ``code`` on the worker's loop, or return ``None`` once closed.""" + with self._lock: + if self._closed: + return None + loop = self._ensure_loop() + # The store is read from the worker's thread while the caller + # waits. A store only grows, and each read is one dict operation. + future = asyncio.run_coroutine_threadsafe( + self._worker.run(code, handles), loop + ) + self._calls.add(future) + future.add_done_callback(self._forget) + return future + + def _forget(self, future: concurrent.futures.Future[_Reply]) -> None: + """Drop a finished call from the set that ``close()`` cancels.""" + with self._lock: + self._calls.discard(future) + + def _ensure_loop(self) -> asyncio.AbstractEventLoop: + """Return the loop, starting it and its thread if needed. Needs the lock.""" + if self._loop is None: + loop = asyncio.new_event_loop() + thread = threading.Thread( + target=_run_loop, + args=(loop,), + name="commons-python-session", + daemon=True, + ) + thread.start() + self._loop, self._thread = loop, thread + return self._loop + + +def _run_loop(loop: asyncio.AbstractEventLoop) -> None: + """Run ``loop`` until it is stopped, then close it, however close() was called.""" + try: + loop.run_forever() + finally: + loop.close() diff --git a/pkg-py/src/commons/_run_python.py b/pkg-py/src/commons/_run_python.py new file mode 100644 index 00000000..716a4bc2 --- /dev/null +++ b/pkg-py/src/commons/_run_python.py @@ -0,0 +1,422 @@ +"""The `run_python` tool, which runs the model's code in a sandboxed session. + +`pkg-r/R/run-r.R` builds R's `run_r`. The two tools follow one contract for +what they tell the model and what they return, each in its own language's +idiom. + +The agent registers the async version of the tool, which waits for the session +without blocking the caller's event loop. Because chatlas refuses a synchronous +`chat()` while any async tool is registered, `Commons.chat()` dynamically swaps the +sync version for the length of the call if required. +""" + +from __future__ import annotations + +import base64 +import html +import io +import keyword +import subprocess +import sys +import tempfile +import tokenize +from collections.abc import Sequence +from typing import Any + +from chatlas import ContentToolResult, Tool +from chatlas.types import ContentImageInline, ContentText +from htmltools import HTML, Tag, div, tags + +from ._citations import tool_result +from ._display import CODE_ANALYSIS, visible_result_note +from ._execution._backend import Network +from ._execution._driver import Failure +from ._execution._env import worker_env +from ._execution._protocol import Error, OpaqueValue, Plot, Result, Text +from ._execution._sandbox import needs_single_thread +from ._execution._thread import WorkerThread +from ._frames import describe_frame, is_frame +from ._prompt import EXECUTION_TOOL +from ._provenance import Tag as ProvenanceTag +from ._rows import frame_rows, rows_to_markdown +from ._tools import ToolContext + +__all__ = [ + "HANDLE_TOOLS", + "run_python_description", + "run_python_result", + "run_python_tools", +] + +# The tools whose results are stored as handles, in the order the +# description names them. +HANDLE_TOOLS = ("call_measure", "call_metrics", "run_sql") + +NO_OUTPUT = "(The code ran but produced no output.)" + +RESTART_NOTE = ( + "(A previous call was cancelled, so the Python session was restarted " + "before this call. Session variables were reset.)" +) + +# How long to wait for the session's interpreter to say what it can import. +PROBE_TIMEOUT = 10.0 + + +def run_python_description( + tool_names: Sequence[str], + *, + has_measures: bool, + network: Network, + can_plot: bool | None = None, + can_install: bool | None = None, +) -> str: + """The description that tells the model what `run_python` does. + + ``tool_names`` are the agent's other registered tools. The description + lists the results of those that store them as preloaded variables. + ``can_plot`` and ``can_install`` say whether the session can import + matplotlib and pip. When omitted, they are found by asking the session's + interpreter; pip is checked only when the session has network access. + """ + if can_plot is None: + can_plot = session_can_import("matplotlib") + if can_install is None and network != "none": + can_install = session_can_import("pip") + handle_tools = [name for name in HANDLE_TOOLS if name in tool_names] + parts = [ + ( + "Run Python code in your sandboxed Python session to analyze results or " + "render plots. Python code and textual output are visible only to you; " + "rendered plots are also shown to the user." + if can_plot + else "Run Python code in your sandboxed Python session to analyze " + "results. Python code and its output are visible only to you." + ), + ( + "The user cannot access or interact with this session. Never direct them " + "to run code or inspect its variables or files; perform follow-up " + "analysis yourself and report the result in your response." + ), + ( + "Your session persists across calls: variables you assign and modules " + "you import remain available." + ), + ] + if handle_tools: + parts.append( + f"Results from {_listed(handle_tools)} are preloaded as variables " + "(r1, r2, ...)." + ) + if has_measures: + parts.append( + "Measure definitions and their helper functions are predefined under " + "their own names: call inspect.getsource() on a measure to read its " + "source. These are source-only copies without their original " + "environment or database connections, so treat them as reference " + "material; to compute a measure, use call_measure." + ) + rules = [ + "Work incrementally: each call should do one small, well-defined task.", + ( + "Follow PEP 8: put separate statements on separate lines and wrap long " + "calls for readability." + ), + ( + "Create at most one matplotlib figure per call. Draw it with pyplot " + "and do not call savefig, since a saved file reaches neither you nor " + "the user. The figure appears where you call plt.show() or " + "fig.show(), or after the call's text if you call neither." + if can_plot + else "matplotlib is not installed, so the session cannot draw plots." + ), + ( + "Do not use this tool to talk to the user; explanations belong in your " + "reply." + ), + ( + "Return results by ending with an expression (`df`, not `print(df)`) " + "and prefer brief summaries (df.head(), df.describe()) over large " + "outputs." + ), + "The session can only write to its own temporary directory.", + ] + if network == "none": + rules.append("The session has no network access.") + elif can_install: + rules.append( + "The session has network access. The sandbox stops subprocesses, so " + "to use a package that is not installed, run pip in the session: " + "`import sys, tempfile; from pip._internal.cli.main import main; " + "path = tempfile.mkdtemp(); main(['install', '--no-cache-dir', " + "'--target', path, 'PACKAGE']); sys.path.insert(0, path)`, with " + "PACKAGE replaced by the package's name." + ) + else: + rules.append( + "The session has network access, but only packages that are " + "already installed can be imported." + ) + return " ".join(parts) + "\n\nRules:" + "".join(f"\n- {rule}" for rule in rules) + + +# Definite answers from the probe, by module. A probe that failed to run is +# not kept, so a slow first start does not settle the answer for the process. +_IMPORTABLE: dict[str, bool] = {} + + +def session_can_import(module: str) -> bool: + """Whether the session's interpreter can import the top-level ``module``. + + The session starts Python in isolated mode (``-I``), which ignores the + user's site-packages directory, ``PYTHONPATH``, and the current directory. + The check therefore starts the same interpreter with the same ``-I`` + isolation, the session's environment variables, and an empty temporary + directory, though outside the sandbox. A definite answer is cached, + because the installed packages do not change while the process runs; a + probe that fails to run counts as no, and is asked again next time. + """ + if module in _IMPORTABLE: + return _IMPORTABLE[module] + probe = "import importlib.util, sys; sys.exit(importlib.util.find_spec(sys.argv[1]) is None)" + try: + with tempfile.TemporaryDirectory(prefix="commons-probe-") as scratch: + completed = subprocess.run( + [sys.executable, "-I", "-c", probe, module], + capture_output=True, + cwd=scratch, + env=worker_env(scratch, single_thread=needs_single_thread()), + timeout=PROBE_TIMEOUT, + check=False, + ) + except (OSError, subprocess.SubprocessError): + return False + _IMPORTABLE[module] = completed.returncode == 0 + return _IMPORTABLE[module] + + +def _listed(names: Sequence[str]) -> str: + if len(names) <= 2: + return " and ".join(names) + return f"{', '.join(names[:-1])}, and {names[-1]}" + + +def run_python_tools( + runner: WorkerThread, context: ToolContext, description: str, network: Network +) -> tuple[Tool, Tool]: + """The async tool the agent registers, and the sync tool `chat()` uses instead.""" + + def finish(code: str, reply: Result | Error | Failure) -> ContentToolResult: + result = run_python_result(code, reply, restarted=runner.take_restart()) + if context.citation_request is None: + return result + return context.citation_request.add_request(result) + + async def run_python(code: str) -> ContentToolResult: + return finish(code, await runner.run(code, context.handles)) + + def run_python_sync(code: str) -> ContentToolResult: + return finish(code, runner.run_sync(code, context.handles)) + + parameters = { + "type": "object", + "properties": { + "code": {"type": "string", "description": "The Python code to run."} + }, + "required": ["code"], + "additionalProperties": False, + } + annotations: Any = { + "title": CODE_ANALYSIS.running, + "readOnlyHint": False, + "openWorldHint": network == "full", + } + + def build(func: Any) -> Tool: + return Tool( + func=func, + name=EXECUTION_TOOL, + description=description, + parameters=parameters, + annotations=annotations, + ) + + return build(run_python), build(run_python_sync) + + +def run_python_result( + code: str, reply: Result | Error | Failure, *, restarted: bool = False +) -> ContentToolResult: + """The tool result for one call: one view for the model, one for the user. + + The model gets the call's output, with each plot as an image at the point + where it was drawn. The user sees the code and its output, followed by the + plots. ``restarted`` says a cancelled call restarted the session before + this one, and the model is told so before the output. + """ + runs = _runs(reply) + plots = [run for run in runs if isinstance(run, Plot)] + return tool_result( + _model_value([RESTART_NOTE, *runs] if restarted else runs), + ProvenanceTag.B, + title=CODE_ANALYSIS.settled, + html=_display_html(code, runs), + open=bool(plots), + ) + + +def _runs(reply: Result | Error | Failure) -> list[str | Plot]: + """The reply's output in order, with consecutive text joined into one string. + + Text written to stdout and stderr is joined in the order it was written, so + a line printed in two parts stays one line. The call's final value, or its + error, goes on a new line after all the other output. + """ + if isinstance(reply, Failure): + return [f"Error: {reply.message}"] + if isinstance(reply, Error): + last = reply.traceback.strip() or f"Error: {reply.message}" + else: + last = _value_text(reply.value) + runs: list[str | Plot] = [] + text = "" + for segment in reply.output: + if isinstance(segment, Text): + text += segment.text + continue + if text.strip("\n"): + runs.append(text.rstrip("\n")) + text = "" + runs.append(segment) + if last: + text += ("\n" if text and not text.endswith("\n") else "") + last + if text.strip("\n"): + runs.append(text.rstrip("\n")) + return runs + + +def _value_text(value: Any) -> str: + """The call's final value as the Python REPL shows it; empty for None.""" + if value is None: + return "" + if isinstance(value, OpaqueValue): + return value.text + if is_frame(value): + rows = frame_rows(value) + return rows_to_markdown(rows) if rows is not None else describe_frame(value) + return repr(value) + + +def _model_value(runs: list[str | Plot]) -> Any: + if not any(isinstance(run, Plot) for run in runs): + return "\n".join(run for run in runs if isinstance(run, str)) or NO_OUTPUT + parts: list[Any] = [ + ContentImageInline( + image_content_type="image/png", + data=base64.b64encode(run.png).decode("ascii"), + ) + if isinstance(run, Plot) + else ContentText(text=run) + for run in runs + ] + parts.append(ContentText(text=visible_result_note("plot"))) + return parts + + +def _display_html(code: str, runs: list[str | Plot]) -> Tag: + output = [ + f"#> {line}" for run in runs if isinstance(run, str) for line in run.split("\n") + ] + plots = [run for run in runs if isinstance(run, Plot)] + block: Tag = tags.pre( + tags.code( + HTML(highlight_python("\n".join([code, *output]))), + class_="language-python", + ), + class_="commons-run-code", + ) + if plots: + block = tags.details( + tags.summary("Details"), block, class_="commons-run-details" + ) + images = [ + tags.img( + class_="commons-run-plot", + src="data:image/png;base64," + + base64.b64encode(plot.display_png).decode("ascii"), + alt="Plot produced by Python code", + width=str(plot.width), + height=str(plot.height), + ) + for plot in plots + ] + return div(block, *images, class_="commons-run-display") + + +# The highlight classes the shared stylesheet styles, by token kind. +_COMMENT, _KEYWORD, _CALL, _NUMBER, _STRING = "com", "kwa", "kwd", "num", "sng" +_STRING_TOKENS = { + name + for name in ("STRING", "FSTRING_START", "FSTRING_MIDDLE", "FSTRING_END") + if hasattr(tokenize, name) +} + + +def highlight_python(source: str) -> str: + """``source`` as escaped HTML, with each token in a span the stylesheet colors. + + Text that Python cannot tokenize is escaped without highlighting. + """ + # Split as tokenize reads, on newlines only, so offsets line up with its + # rows even when the text holds a carriage return or a form feed. + lines = io.StringIO(source).readlines() + starts = [0] + for line in lines: + starts.append(starts[-1] + len(line)) + + def offset(position: tuple[int, int]) -> int: + row, column = position + return starts[row - 1] + column if row <= len(lines) else len(source) + + spans: list[tuple[int, int, str]] = [] + try: + tokens = list(tokenize.generate_tokens(io.StringIO(source).readline)) + except (tokenize.TokenError, SyntaxError, UnicodeDecodeError): + # Python 3.12+ raises UnicodeDecodeError for a lone carriage return + # before non-ASCII text. + return html.escape(source) + for index, token in enumerate(tokens): + kind = _token_class( + token, tokens[index + 1] if index + 1 < len(tokens) else None + ) + if kind is not None: + spans.append((offset(token.start), offset(token.end), kind)) + + out: list[str] = [] + cursor = 0 + for start, end, kind in spans: + if start < cursor: + continue + out.append(html.escape(source[cursor:start])) + out.append(f'{html.escape(source[start:end])}') + cursor = end + out.append(html.escape(source[cursor:])) + return "".join(out) + + +def _token_class( + token: tokenize.TokenInfo, following: tokenize.TokenInfo | None +) -> str | None: + name = tokenize.tok_name[token.type] + if token.type == tokenize.COMMENT: + return _COMMENT + if token.type == tokenize.NUMBER: + return _NUMBER + if name in _STRING_TOKENS: + return _STRING + if token.type == tokenize.NAME: + if keyword.iskeyword(token.string): + return _KEYWORD + if following is not None and following.string == "(": + return _CALL + return None diff --git a/pkg-py/src/commons/www/commons-chat/commons-chat.css b/pkg-py/src/commons/www/commons-chat/commons-chat.css index 14af90dd..fb8f4647 100644 --- a/pkg-py/src/commons/www/commons-chat/commons-chat.css +++ b/pkg-py/src/commons/www/commons-chat/commons-chat.css @@ -556,14 +556,14 @@ shiny-chat-container white-space: normal; } -/* ---- run_r display ---------------------------------------------------- */ +/* ---- run_r and run_python display -------------------------------------- */ -.commons-run-r-display { +.commons-run-display { display: grid; gap: 0.6rem; } -.commons-run-r-details > summary { +.commons-run-details > summary { align-items: center; color: inherit; cursor: pointer; @@ -574,11 +574,11 @@ shiny-chat-container width: fit-content; } -.commons-run-r-details > summary::-webkit-details-marker { +.commons-run-details > summary::-webkit-details-marker { display: none; } -.commons-run-r-details > summary::before { +.commons-run-details > summary::before { border-bottom: 1px solid currentColor; border-right: 1px solid currentColor; content: ""; @@ -590,90 +590,90 @@ shiny-chat-container width: 0.45em; } -.commons-run-r-details[open] > summary::before { +.commons-run-details[open] > summary::before { transform: rotate(45deg); } -.commons-run-r-details > summary:focus-visible { +.commons-run-details > summary:focus-visible { border-radius: var(--bs-border-radius-sm, 0.25rem); outline: 2px solid var(--bs-primary, #007bc2); outline-offset: 2px; } -.commons-run-r-details[open] > .commons-run-r-code { - animation: commons-run-r-details-reveal 0.2s ease-out; +.commons-run-details[open] > .commons-run-code { + animation: commons-run-details-reveal 0.2s ease-out; } -@keyframes commons-run-r-details-reveal { +@keyframes commons-run-details-reveal { from { opacity: 0; } } @media (prefers-reduced-motion: reduce) { - .commons-run-r-details > summary::before { + .commons-run-details > summary::before { transition: none; } - .commons-run-r-details[open] > .commons-run-r-code { + .commons-run-details[open] > .commons-run-code { animation: none; } } -.commons-run-r-display .commons-run-r-code { +.commons-run-display .commons-run-code { margin: 0; overflow-x: auto; white-space: pre; } -.commons-run-r-code .hl.com { +.commons-run-code .hl.com { color: #a0a1a7; font-style: italic; } -.commons-run-r-code .hl.kwa { +.commons-run-code .hl.kwa { color: #a626a4; } -.commons-run-r-code .hl.kwc, -.commons-run-r-code .hl.num { +.commons-run-code .hl.kwc, +.commons-run-code .hl.num { color: #986801; } -.commons-run-r-code .hl.kwd { +.commons-run-code .hl.kwd { color: #4078f2; } -.commons-run-r-code .hl.sng { +.commons-run-code .hl.sng { color: #50a14f; } -[data-bs-theme="dark"] .commons-run-r-code .hl.com { +[data-bs-theme="dark"] .commons-run-code .hl.com { color: #5c6370; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwa { +[data-bs-theme="dark"] .commons-run-code .hl.kwa { color: #c678dd; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwc, -[data-bs-theme="dark"] .commons-run-r-code .hl.num { +[data-bs-theme="dark"] .commons-run-code .hl.kwc, +[data-bs-theme="dark"] .commons-run-code .hl.num { color: #d19a66; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwd { +[data-bs-theme="dark"] .commons-run-code .hl.kwd { color: #61aeee; } -[data-bs-theme="dark"] .commons-run-r-code .hl.sng { +[data-bs-theme="dark"] .commons-run-code .hl.sng { color: #98c379; } -.commons-run-r-details > .commons-run-r-code { +.commons-run-details > .commons-run-code { margin-top: 0.4rem; } -.commons-run-r-plot { +.commons-run-plot { border: 1px solid var(--bs-border-color, #dee2e6); border-radius: 0.5rem; height: auto; diff --git a/pkg-py/tests/test_agent.py b/pkg-py/tests/test_agent.py index 295fb317..deca85fc 100644 --- a/pkg-py/tests/test_agent.py +++ b/pkg-py/tests/test_agent.py @@ -345,6 +345,7 @@ def test_an_agent_registers_the_tools_its_composition_earns( "describe_table", "run_sql", "load_skill", + "run_python", ] diff --git a/pkg-py/tests/test_execution_thread.py b/pkg-py/tests/test_execution_thread.py new file mode 100644 index 00000000..58d9782d --- /dev/null +++ b/pkg-py/tests/test_execution_thread.py @@ -0,0 +1,224 @@ +"""The worker thread: one session, reachable from sync callers and any event loop.""" + +from __future__ import annotations + +import asyncio +import signal +import threading +import time + +import pytest + +from commons._execution._driver import Failure, Worker +from commons._execution._protocol import Result +from commons._execution._thread import WorkerThread + + +@pytest.fixture +def runner(): + worker_thread = WorkerThread(Worker(call_timeout=10)) + yield worker_thread + worker_thread.close() + + +def test_sync_and_async_callers_share_one_session(runner: WorkerThread) -> None: + assert isinstance(runner.run_sync("x = 41"), Result) + + async def from_a_loop() -> Result: + reply = await runner.run("x + 1") + assert isinstance(reply, Result) + return reply + + assert asyncio.run(from_a_loop()).value == 42 + + +def test_the_session_runs_off_the_callers_loop(runner: WorkerThread) -> None: + async def caller() -> tuple[int, bool]: + ticks = 0 + stop = asyncio.Event() + + async def tick() -> None: + nonlocal ticks + while not stop.is_set(): + ticks += 1 + await asyncio.sleep(0.01) + + ticker = asyncio.ensure_future(tick()) + reply = await runner.run("import time; time.sleep(0.5)") + stop.set() + await ticker + return ticks, isinstance(reply, Result) + + ticks, ok = asyncio.run(caller()) + assert ok + # The caller's loop kept running while the call slept. + assert ticks > 10 + + +def test_no_thread_starts_before_the_first_call() -> None: + worker_thread = WorkerThread(Worker()) + assert worker_thread._thread is None + worker_thread.close() + assert worker_thread._thread is None + + +def test_run_sync_refuses_the_workers_own_thread(runner: WorkerThread) -> None: + runner.run_sync("1") + assert runner._loop is not None + + async def from_the_loop() -> None: + runner.run_sync("1") + + future = asyncio.run_coroutine_threadsafe(from_the_loop(), runner._loop) + with pytest.raises(RuntimeError, match="deadlock"): + future.result(10) + + +def test_ctrl_c_while_waiting_stops_the_call() -> None: + worker_thread = WorkerThread(Worker(call_timeout=60)) + main = threading.main_thread().ident + assert main is not None + interrupt = threading.Timer(1, signal.pthread_kill, (main, signal.SIGINT)) + try: + worker_thread.run_sync("1") + interrupt.start() + with pytest.raises(KeyboardInterrupt): + worker_thread.run_sync("import time; time.sleep(60)") + started = time.monotonic() + # Without the cancel, this call would queue behind the sleep. + reply = worker_thread.run_sync("1 + 1") + assert time.monotonic() - started < 15 + assert isinstance(reply, Result) + assert reply.value == 2 + assert worker_thread.take_restart() + assert not worker_thread.take_restart() + finally: + # A call that returned early must not leave the signal for a later test. + interrupt.cancel() + worker_thread.close() + + +def test_close_from_the_loops_own_thread_returns_and_the_thread_stops() -> None: + worker_thread = WorkerThread(Worker()) + worker_thread.run_sync("1") + loop, thread = worker_thread._loop, worker_thread._thread + assert loop is not None and thread is not None + loop.call_soon_threadsafe(worker_thread.close) + thread.join(15) + assert not thread.is_alive() + assert loop.is_closed() + + +def test_a_cancelled_caller_cancels_the_call() -> None: + worker_thread = WorkerThread(Worker(call_timeout=60)) + + async def cancel_midway() -> None: + call = asyncio.ensure_future(worker_thread.run("import time; time.sleep(60)")) + await asyncio.sleep(1) + call.cancel() + with pytest.raises(asyncio.CancelledError): + await call + + try: + asyncio.run(cancel_midway()) + started = time.monotonic() + # The cancelled call's worker was shut down, so the next one respawns + # rather than queue behind the sleep. + reply = worker_thread.run_sync("1 + 1") + assert time.monotonic() - started < 15 + assert isinstance(reply, Result) + assert reply.value == 2 + assert worker_thread.take_restart() + finally: + worker_thread.close() + + +def test_close_ends_the_thread_and_later_calls_fail(runner: WorkerThread) -> None: + runner.run_sync("1") + runner.close() + runner.close() + reply = runner.run_sync("1") + assert isinstance(reply, Failure) + assert "closed" in reply.message + + +def test_close_cancels_a_running_call_rather_than_waiting_it_out() -> None: + worker_thread = WorkerThread(Worker(call_timeout=60)) + try: + worker_thread.run_sync("1") + replies: list[object] = [] + caller = threading.Thread( + target=lambda: replies.append( + worker_thread.run_sync("import time; time.sleep(60)") + ) + ) + caller.start() + time.sleep(1) + started = time.monotonic() + worker_thread.close() + caller.join(10) + assert time.monotonic() - started < 15 + assert replies == [Failure(message="the Python session is closed.")] + finally: + worker_thread.close() + + +def test_an_async_call_running_when_the_thread_closes_is_told_it_closed() -> None: + worker_thread = WorkerThread(Worker(call_timeout=60)) + + async def caller() -> object: + await worker_thread.run("1") + call = asyncio.ensure_future(worker_thread.run("import time; time.sleep(60)")) + await asyncio.sleep(1) + # close() blocks, so it runs off this loop, which keeps serving the call. + await asyncio.to_thread(worker_thread.close) + return await asyncio.wait_for(call, 10) + + try: + assert asyncio.run(caller()) == Failure( + message="the Python session is closed." + ) + finally: + worker_thread.close() + + +class FailingClose(Worker): + async def aclose(self) -> None: + await super().aclose() + raise OSError("shutdown failed") + + +def test_a_shutdown_that_raises_still_stops_the_thread_and_closes_the_loop() -> None: + worker_thread = WorkerThread(FailingClose()) + worker_thread.run_sync("1") + loop, thread = worker_thread._loop, worker_thread._thread + assert loop is not None and thread is not None + with pytest.raises(OSError, match="shutdown failed"): + worker_thread.close() + assert not thread.is_alive() + assert loop.is_closed() + + +def test_a_call_cancelled_while_it_waits_leaves_the_session_alone() -> None: + worker_thread = WorkerThread(Worker(call_timeout=60)) + + async def cancel_the_queued_one() -> object: + await worker_thread.run("x = 41") + running = asyncio.ensure_future( + worker_thread.run("import time; time.sleep(2)\nx + 1") + ) + await asyncio.sleep(0.5) + queued = asyncio.ensure_future(worker_thread.run("1")) + await asyncio.sleep(0.5) + queued.cancel() + with pytest.raises(asyncio.CancelledError): + await queued + return await running + + try: + reply = asyncio.run(cancel_the_queued_one()) + assert isinstance(reply, Result) + assert reply.value == 42 + assert not worker_thread.take_restart() + finally: + worker_thread.close() diff --git a/pkg-py/tests/test_run_python.py b/pkg-py/tests/test_run_python.py new file mode 100644 index 00000000..4a6cbe03 --- /dev/null +++ b/pkg-py/tests/test_run_python.py @@ -0,0 +1,420 @@ +"""run_python: its description, the result it builds, and calls through an agent.""" + +from __future__ import annotations + +import base64 +import gc +import html +import importlib.util +import re +from pathlib import Path +from typing import Any, NoReturn + +import pandas as pd +import pytest +from chatlas import ContentToolRequest, ContentToolResult, Tool +from chatlas.types import ContentImageInline, ContentText + +from commons import Injected, data_source, measure, semantic_layer +from commons._agent import Commons +from commons._display import DISPLAY_EXTRA_KEY +from commons._execution._driver import Failure +from commons._execution._protocol import Error, OpaqueValue, Plot, Result, Text +from commons._provenance import TAG_EXTRA_KEY, Tag +from commons._run_python import ( + NO_OUTPUT, + RESTART_NOTE, + highlight_python, + run_python_description, + run_python_result, + session_can_import, +) + +from ._provider import scripted_chat, text + + +@measure(description="Total revenue.") +def total_revenue(sales: Injected[Any]) -> float: + return float(sales.execute("SELECT SUM(revenue) FROM sales").fetchone()[0]) + + +def frame() -> pd.DataFrame: + return pd.DataFrame({"revenue": [500.0, 900.0, 300.0]}) + + +def png(width: int, height: int) -> bytes: + """A PNG signature and IHDR header of the given size, which is all Plot reads.""" + header = width.to_bytes(4, "big") + height.to_bytes(4, "big") + return ( + b"\x89PNG\r\n\x1a\n" + b"\x00\x00\x00\rIHDR" + header + b"\x08\x06\x00\x00\x00" + ) + + +def plot() -> Plot: + return Plot(png=png(400, 300), display_png=png(800, 600)) + + +def display(result: ContentToolResult) -> dict[str, Any]: + assert result.extra is not None + return result.extra[DISPLAY_EXTRA_KEY] + + +def tool_request(**arguments: Any) -> list[Any]: + return [ + ContentToolRequest(id="call-run_python", name="run_python", arguments=arguments) + ] + + +def run_python_tool(agent: Commons) -> Tool: + tool = {tool.name: tool for tool in agent.get_tools()}["run_python"] + assert isinstance(tool, Tool) + return tool + + +def tool_results(agent: Commons) -> list[ContentToolResult]: + return [ + content + for turn in agent.get_turns() + for content in turn.contents + if isinstance(content, ContentToolResult) + ] + + +# ---- the description -------------------------------------------------------- + + +def test_the_preloaded_handles_name_the_registered_tools_that_store_results() -> None: + description = run_python_description( + ["search_pool", "call_measure", "run_sql"], has_measures=True, network="none" + ) + assert "Results from call_measure and run_sql are preloaded" in description + description = run_python_description( + ["call_measure", "call_metrics", "run_sql"], has_measures=True, network="none" + ) + assert "Results from call_measure, call_metrics, and run_sql are preloaded" in ( + description + ) + description = run_python_description( + ["run_sql"], has_measures=False, network="none" + ) + assert "Results from run_sql are preloaded" in description + + +def test_measure_sources_are_described_only_when_there_are_measures() -> None: + with_measures = run_python_description( + ["run_sql"], has_measures=True, network="none" + ) + without = run_python_description(["run_sql"], has_measures=False, network="none") + assert "inspect.getsource()" in with_measures + assert "inspect.getsource()" not in without + + +@pytest.mark.parametrize( + ("network", "can_install", "rule"), + [ + ("none", True, "- The session has no network access."), + ("full", True, "from pip._internal.cli.main import main"), + ("full", False, "only packages that are already installed can be imported"), + ], +) +def test_the_network_rule_says_what_the_session_can_reach( + network: Any, can_install: bool, rule: str +) -> None: + description = run_python_description( + ["run_sql"], has_measures=False, network=network, can_install=can_install + ) + assert rule in description.split("\n\nRules:")[1] + + +def test_the_plot_rule_says_whether_the_session_can_draw() -> None: + def rules(can_plot: bool) -> str: + return run_python_description( + ["run_sql"], has_measures=False, network="none", can_plot=can_plot + ).split("\n\nRules:")[1] + + assert "at most one matplotlib figure per call" in rules(True) + assert "the session cannot draw plots" in rules(False) + + +def test_a_module_the_isolated_session_cannot_see_is_not_importable( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + assert session_can_import("pytest") + assert not session_can_import("no_such_module_for_commons") + # On PYTHONPATH and so on this process's path, but -I ignores both. + (tmp_path / "only_on_pythonpath.py").write_text("") + monkeypatch.setenv("PYTHONPATH", str(tmp_path)) + monkeypatch.syspath_prepend(str(tmp_path)) + assert importlib.util.find_spec("only_on_pythonpath") is not None + assert not session_can_import("only_on_pythonpath") + + +def test_a_probe_that_fails_to_run_is_asked_again( + monkeypatch: pytest.MonkeyPatch, +) -> None: + # Too short for the interpreter to start, so the probe times out. + monkeypatch.setattr("commons._run_python.PROBE_TIMEOUT", 1e-6) + assert not session_can_import("email") + monkeypatch.undo() + assert session_can_import("email") + + +# ---- the result ------------------------------------------------------------- + + +def test_a_result_shows_what_was_printed_and_the_value() -> None: + output = ( + Text(stream="stdout", text="hi\n"), + Text(stream="stderr", text="warn\n"), + Text(stream="stdout", text="partial"), + Text(stream="stdout", text=" line\n"), + ) + result = run_python_result( + "print('hi')\n1 + 1", Result(id="c1", value=2, output=output) + ) + # Streams interleave as written; the value starts its own line. + assert result.value == "hi\nwarn\npartial line\n2" + assert result.extra is not None + assert result.extra[TAG_EXTRA_KEY] == Tag.B + assert display(result)["title"] == "Analyzed data" + assert display(result)["open"] is False + + +def test_a_restarted_session_is_noted_for_the_model_only() -> None: + result = run_python_result("x", Result(id="c1", value=2), restarted=True) + assert result.value == f"{RESTART_NOTE}\n2" + assert "restarted" not in str(display(result)["html"]) + + +def test_a_call_that_shows_nothing_says_so() -> None: + assert run_python_result("x = 1", Result(id="c1")).value == NO_OUTPUT + + +def test_a_frame_value_is_a_table_and_an_opaque_value_is_its_repr() -> None: + table = run_python_result("df", Result(id="c1", value=frame())).value + assert "| revenue |" in table + opaque = OpaqueValue(type_name="Connection", text="") + assert run_python_result("con", Result(id="c1", value=opaque)).value == ( + "" + ) + + +def test_an_error_shows_the_traceback_and_a_failure_its_message() -> None: + error = Error( + id="c1", + message="NameError: name 'y' is not defined", + traceback="Traceback...\nNameError", + output=(Text(stream="stdout", text="got this far"),), + ) + # What ran before the error is kept, so the model sees how far it got. + assert run_python_result("y", error).value == ( + "got this far\nTraceback...\nNameError" + ) + failure = run_python_result("1", Failure(message="the Python session crashed.")) + assert failure.value == "Error: the Python session crashed." + assert failure.extra is not None + assert failure.extra[TAG_EXTRA_KEY] == Tag.B + + +def test_plots_reach_the_model_as_images_and_the_reader_at_their_size() -> None: + output = ( + Text(stream="stdout", text="before\n"), + plot(), + Text(stream="stdout", text="after\n"), + ) + result = run_python_result("fig", Result(id="c1", value=3, output=output)) + parts = result.value + assert isinstance(parts, list) + # The plot keeps its place between the text written before and after it. + assert [type(part) for part in parts] == [ + ContentText, + ContentImageInline, + ContentText, + ContentText, + ] + assert parts[0].text == "before" + assert base64.b64decode(parts[1].data) == plot().png + assert parts[2].text == "after\n3" + assert "already visible to the user" in parts[-1].text + shown = display(result) + assert shown["open"] is True + html = str(shown["html"]) + assert 'class="commons-run-details"' in html + assert base64.b64encode(plot().display_png).decode() in html + assert 'width="400"' in html and 'height="300"' in html + + +def test_the_display_shows_the_code_and_its_output_escaped() -> None: + html = str( + display(run_python_result("''", Result(id="c1", value="")))["html"] + ) + assert 'class="commons-run-code"' in html + assert "#> '<b>'" in html or "#> '<b>'" in html + assert "" not in html + + +# ---- highlighting ----------------------------------------------------------- + + +def test_highlighting_marks_tokens_with_the_stylesheets_classes() -> None: + out = highlight_python("def f(x):\n return len('a') + 1 # note\n") + assert 'def' in out + assert 'len' in out + assert ''a'' in out + assert '1' in out + assert '# note' in out + + +def test_text_that_does_not_tokenize_is_escaped_plainly() -> None: + assert highlight_python("''' None: + # Python 3.12+ tokenize raises UnicodeDecodeError on this input, and + # earlier versions highlight it, so only the text is the same everywhere. + source = "#> 50%\r\u00e9t\u00e9" + assert html.unescape(re.sub(r"<[^>]+>", "", highlight_python(source))) == source + + +def test_a_carriage_return_or_form_feed_does_not_shift_the_highlighting() -> None: + out = highlight_python('#> 50%\r100%\n# a\x0cb\ng("z")') + assert 'g("z")' in out + assert html.unescape(re.sub(r"<[^>]+>", "", out)) == '#> 50%\r100%\n# a\x0cb\ng("z")' + + +# ---- through an agent ------------------------------------------------------- + + +def test_an_agent_registers_run_python_with_the_measure_source_note() -> None: + agent = Commons( + scripted_chat(), + {"sales": data_source(sales=frame())}, + semantic_layer(total_revenue), + ) + description = run_python_tool(agent).schema["function"]["description"] + assert "inspect.getsource()" in description + assert "Results from call_measure and run_sql are preloaded" in description + + +def test_chat_runs_code_against_a_preloaded_handle() -> None: + agent = Commons( + scripted_chat( + [ + [ + ContentToolRequest( + id="q", name="run_sql", arguments={"sql": "SELECT * FROM sales"} + ) + ], + tool_request(code="sum(row['revenue'] for row in r1)"), + text("Done."), + ] + ), + data_source(sales=frame()), + ) + agent.chat("Total revenue?", echo="none") + run = tool_results(agent)[-1] + assert run.value.startswith("1700.0") + # The async tool is back in place for stream_async. + assert run_python_tool(agent)._is_async + + +async def test_stream_async_keeps_session_state_and_reads_measure_sources() -> None: + agent = Commons( + scripted_chat( + [ + tool_request(code="x = 41"), + tool_request( + code="import inspect\nprint(inspect.getsource(total_revenue))\nx + 1" + ), + text("Done."), + ] + ), + {"sales": data_source(sales=frame())}, + semantic_layer(total_revenue), + ) + stream = await agent.stream_async("Go.") + [chunk async for chunk in stream] + run = tool_results(agent)[-1] + assert "def total_revenue(sales" in run.value + assert "\n42" in run.value + + +@pytest.mark.skipif( + importlib.util.find_spec("matplotlib") is None, reason="needs matplotlib" +) +async def test_a_plot_reaches_the_model_as_an_image() -> None: + provider_chat = scripted_chat( + [ + tool_request(code="import matplotlib.pyplot as plt\nplt.plot([1, 2, 3])"), + text("Done."), + ] + ) + agent = Commons(provider_chat, data_source(sales=frame())) + stream = await agent.stream_async("Plot it.") + [chunk async for chunk in stream] + run = tool_results(agent)[-1] + assert any(isinstance(part, ContentImageInline) for part in run.value) + # chatlas moves the image out of the tool result for the provider. + last_request = agent.provider.requests[-1] # type: ignore[attr-defined] + sent = [content for turn in last_request for content in turn.contents] + assert any(isinstance(content, ContentImageInline) for content in sent) + + +def test_the_network_argument_reaches_the_description_and_annotations() -> None: + agent = Commons(scripted_chat(), data_source(sales=frame()), network="full") + tool = run_python_tool(agent) + rules = tool.schema["function"]["description"].split("\n\nRules:")[1] + assert "The session has network access" in rules + assert tool.annotations is not None + assert tool.annotations.get("openWorldHint") is True + assert agent._python._worker._network == "full" + + +def test_a_network_other_than_none_or_full_is_refused() -> None: + with pytest.raises(ValueError): + Commons(scripted_chat(), data_source(sales=frame()), network="some") # type: ignore[arg-type] + + +def test_chat_restores_the_async_tool_when_the_turn_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + agent = Commons(scripted_chat(), data_source(sales=frame())) + + def unavailable(**_: Any) -> NoReturn: + raise ConnectionError("provider unavailable") + + monkeypatch.setattr(agent.provider, "chat_perform", unavailable) + with pytest.raises(ConnectionError): + agent.chat("Hello?", echo="none") + assert run_python_tool(agent)._is_async + + +def test_chat_leaves_a_removed_run_python_removed() -> None: + agent = Commons(scripted_chat([text("Hi.")]), data_source(sales=frame())) + agent.set_tools([tool for tool in agent.get_tools() if tool.name != "run_python"]) + agent.chat("Hello?", echo="none") + assert "run_python" not in {tool.name for tool in agent.get_tools()} + + +def test_the_first_run_python_result_of_a_turn_asks_for_citations() -> None: + agent = Commons( + scripted_chat([tool_request(code="1 + 1"), text("Done.")]), + data_source(sales=frame()), + ) + agent.chat("What is one plus one?", echo="none") + run = tool_results(agent)[-1] + assert run.value == f"2\n\n{agent._citation_request.reminder}" + + +def test_dropping_the_agent_stops_its_session_thread() -> None: + agent = Commons( + scripted_chat([tool_request(code="1"), text("Done.")]), + data_source(sales=frame()), + ) + agent.chat("Go.", echo="none") + thread = agent._python._thread + assert thread is not None and thread.is_alive() + del agent + gc.collect() + thread.join(15) + assert not thread.is_alive() diff --git a/pkg-r/NEWS.md b/pkg-r/NEWS.md index 0913d478..9482608a 100644 --- a/pkg-r/NEWS.md +++ b/pkg-r/NEWS.md @@ -2,6 +2,8 @@ * `commons()` gains a `mode` argument. With `mode = "trusted only"`, the agent answers only with trusted calculations. It never writes its own SQL or R, and it tells the user when no trusted calculation answers a question. These answers carry no provenance markers. +* The CSS classes on `run_r`'s display are now `commons-run-display`, `commons-run-details`, `commons-run-code`, and `commons-run-plot`, without the `-r`, because the Python package's code tool uses the same classes. App CSS that targets the old `commons-run-r-*` names needs updating. + # commons 0.1.1 * Fixes an issue with the R code sandbox where the generated policy would be diff --git a/pkg-r/R/run-r.R b/pkg-r/R/run-r.R index 913e879d..20fc2553 100644 --- a/pkg-r/R/run-r.R +++ b/pkg-r/R/run-r.R @@ -226,7 +226,7 @@ run_r_html <- function(code, segments) { dims <- plot_dimensions() plot_html <- c(plot_html, sprintf( paste0( - "" ), plot_image_data(seg$path), @@ -241,18 +241,18 @@ run_r_html <- function(code, segments) { } } code_html <- sprintf( - "
%s
", + "
%s
", highlight_r_html(paste(c(code, output), collapse = "\n")) ) if (length(plot_html)) { code_html <- paste0( - "
Details", + "
Details", code_html, "
" ) } sprintf( - "
%s
", + "
%s
", paste(c(code_html, plot_html), collapse = "\n") ) } diff --git a/pkg-r/inst/www/commons-chat/commons-chat.css b/pkg-r/inst/www/commons-chat/commons-chat.css index 14af90dd..fb8f4647 100644 --- a/pkg-r/inst/www/commons-chat/commons-chat.css +++ b/pkg-r/inst/www/commons-chat/commons-chat.css @@ -556,14 +556,14 @@ shiny-chat-container white-space: normal; } -/* ---- run_r display ---------------------------------------------------- */ +/* ---- run_r and run_python display -------------------------------------- */ -.commons-run-r-display { +.commons-run-display { display: grid; gap: 0.6rem; } -.commons-run-r-details > summary { +.commons-run-details > summary { align-items: center; color: inherit; cursor: pointer; @@ -574,11 +574,11 @@ shiny-chat-container width: fit-content; } -.commons-run-r-details > summary::-webkit-details-marker { +.commons-run-details > summary::-webkit-details-marker { display: none; } -.commons-run-r-details > summary::before { +.commons-run-details > summary::before { border-bottom: 1px solid currentColor; border-right: 1px solid currentColor; content: ""; @@ -590,90 +590,90 @@ shiny-chat-container width: 0.45em; } -.commons-run-r-details[open] > summary::before { +.commons-run-details[open] > summary::before { transform: rotate(45deg); } -.commons-run-r-details > summary:focus-visible { +.commons-run-details > summary:focus-visible { border-radius: var(--bs-border-radius-sm, 0.25rem); outline: 2px solid var(--bs-primary, #007bc2); outline-offset: 2px; } -.commons-run-r-details[open] > .commons-run-r-code { - animation: commons-run-r-details-reveal 0.2s ease-out; +.commons-run-details[open] > .commons-run-code { + animation: commons-run-details-reveal 0.2s ease-out; } -@keyframes commons-run-r-details-reveal { +@keyframes commons-run-details-reveal { from { opacity: 0; } } @media (prefers-reduced-motion: reduce) { - .commons-run-r-details > summary::before { + .commons-run-details > summary::before { transition: none; } - .commons-run-r-details[open] > .commons-run-r-code { + .commons-run-details[open] > .commons-run-code { animation: none; } } -.commons-run-r-display .commons-run-r-code { +.commons-run-display .commons-run-code { margin: 0; overflow-x: auto; white-space: pre; } -.commons-run-r-code .hl.com { +.commons-run-code .hl.com { color: #a0a1a7; font-style: italic; } -.commons-run-r-code .hl.kwa { +.commons-run-code .hl.kwa { color: #a626a4; } -.commons-run-r-code .hl.kwc, -.commons-run-r-code .hl.num { +.commons-run-code .hl.kwc, +.commons-run-code .hl.num { color: #986801; } -.commons-run-r-code .hl.kwd { +.commons-run-code .hl.kwd { color: #4078f2; } -.commons-run-r-code .hl.sng { +.commons-run-code .hl.sng { color: #50a14f; } -[data-bs-theme="dark"] .commons-run-r-code .hl.com { +[data-bs-theme="dark"] .commons-run-code .hl.com { color: #5c6370; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwa { +[data-bs-theme="dark"] .commons-run-code .hl.kwa { color: #c678dd; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwc, -[data-bs-theme="dark"] .commons-run-r-code .hl.num { +[data-bs-theme="dark"] .commons-run-code .hl.kwc, +[data-bs-theme="dark"] .commons-run-code .hl.num { color: #d19a66; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwd { +[data-bs-theme="dark"] .commons-run-code .hl.kwd { color: #61aeee; } -[data-bs-theme="dark"] .commons-run-r-code .hl.sng { +[data-bs-theme="dark"] .commons-run-code .hl.sng { color: #98c379; } -.commons-run-r-details > .commons-run-r-code { +.commons-run-details > .commons-run-code { margin-top: 0.4rem; } -.commons-run-r-plot { +.commons-run-plot { border: 1px solid var(--bs-border-color, #dee2e6); border-radius: 0.5rem; height: auto; diff --git a/pkg-r/tests/testthat/test-run-r.R b/pkg-r/tests/testthat/test-run-r.R index f2640821..e084edba 100644 --- a/pkg-r/tests/testthat/test-run-r.R +++ b/pkg-r/tests/testthat/test-run-r.R @@ -31,7 +31,7 @@ test_that("run_r executes code against stored handles", { expect_false(res@extra$display$open) expect_match( res@extra$display$html, - '
',
+    '
',
     fixed = TRUE
   )
   expect_match(
@@ -145,8 +145,8 @@ test_that("run_r returns plots as images and opens the display", {
     "data:image/png;base64,",
     fixed = TRUE
   )
-  expect_match(res@extra$display$html, "commons-run-r-details")
-  expect_match(res@extra$display$html, "commons-run-r-code", fixed = TRUE)
+  expect_match(res@extra$display$html, "commons-run-details")
+  expect_match(res@extra$display$html, "commons-run-code", fixed = TRUE)
   expect_match(res@extra$display$html, "Details", fixed = TRUE)
 })
 
@@ -165,7 +165,7 @@ test_that("run_r collapses code and output above plots", {
   expect_match(res@value[[1]]@text, "private warning")
   expect_match(
     res@extra$display$html,
-    '
Details', + '
Details', fixed = TRUE ) expect_match(res@extra$display$html, "#> private text", fixed = TRUE) @@ -177,7 +177,7 @@ test_that("run_r collapses code and output above plots", { fixed = TRUE ) expect_lt( - as.integer(regexpr("commons-run-r-details", res@extra$display$html)), + as.integer(regexpr("commons-run-details", res@extra$display$html)), as.integer(regexpr( "data:image/png;base64,", res@extra$display$html, @@ -219,7 +219,7 @@ test_that("run_r surfaces errors from model code without failing the tool", { expect_match(res@value, "Error: boom") expect_false(res@extra$display$open) expect_match(res@extra$display$html, "#> boom", fixed = TRUE) - expect_match(res@extra$display$html, "commons-run-r-code", fixed = TRUE) + expect_match(res@extra$display$html, "commons-run-code", fixed = TRUE) expect_no_match(res@extra$display$html, "1', fixed = TRUE ) - expect_match(res@extra$display$html, "commons-run-r-code", fixed = TRUE) + expect_match(res@extra$display$html, "commons-run-code", fixed = TRUE) expect_no_match(res@extra$display$html, "#>", fixed = TRUE) expect_no_match(res@extra$display$html, " summary { +.commons-run-details > summary { align-items: center; color: inherit; cursor: pointer; @@ -574,11 +574,11 @@ shiny-chat-container width: fit-content; } -.commons-run-r-details > summary::-webkit-details-marker { +.commons-run-details > summary::-webkit-details-marker { display: none; } -.commons-run-r-details > summary::before { +.commons-run-details > summary::before { border-bottom: 1px solid currentColor; border-right: 1px solid currentColor; content: ""; @@ -590,90 +590,90 @@ shiny-chat-container width: 0.45em; } -.commons-run-r-details[open] > summary::before { +.commons-run-details[open] > summary::before { transform: rotate(45deg); } -.commons-run-r-details > summary:focus-visible { +.commons-run-details > summary:focus-visible { border-radius: var(--bs-border-radius-sm, 0.25rem); outline: 2px solid var(--bs-primary, #007bc2); outline-offset: 2px; } -.commons-run-r-details[open] > .commons-run-r-code { - animation: commons-run-r-details-reveal 0.2s ease-out; +.commons-run-details[open] > .commons-run-code { + animation: commons-run-details-reveal 0.2s ease-out; } -@keyframes commons-run-r-details-reveal { +@keyframes commons-run-details-reveal { from { opacity: 0; } } @media (prefers-reduced-motion: reduce) { - .commons-run-r-details > summary::before { + .commons-run-details > summary::before { transition: none; } - .commons-run-r-details[open] > .commons-run-r-code { + .commons-run-details[open] > .commons-run-code { animation: none; } } -.commons-run-r-display .commons-run-r-code { +.commons-run-display .commons-run-code { margin: 0; overflow-x: auto; white-space: pre; } -.commons-run-r-code .hl.com { +.commons-run-code .hl.com { color: #a0a1a7; font-style: italic; } -.commons-run-r-code .hl.kwa { +.commons-run-code .hl.kwa { color: #a626a4; } -.commons-run-r-code .hl.kwc, -.commons-run-r-code .hl.num { +.commons-run-code .hl.kwc, +.commons-run-code .hl.num { color: #986801; } -.commons-run-r-code .hl.kwd { +.commons-run-code .hl.kwd { color: #4078f2; } -.commons-run-r-code .hl.sng { +.commons-run-code .hl.sng { color: #50a14f; } -[data-bs-theme="dark"] .commons-run-r-code .hl.com { +[data-bs-theme="dark"] .commons-run-code .hl.com { color: #5c6370; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwa { +[data-bs-theme="dark"] .commons-run-code .hl.kwa { color: #c678dd; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwc, -[data-bs-theme="dark"] .commons-run-r-code .hl.num { +[data-bs-theme="dark"] .commons-run-code .hl.kwc, +[data-bs-theme="dark"] .commons-run-code .hl.num { color: #d19a66; } -[data-bs-theme="dark"] .commons-run-r-code .hl.kwd { +[data-bs-theme="dark"] .commons-run-code .hl.kwd { color: #61aeee; } -[data-bs-theme="dark"] .commons-run-r-code .hl.sng { +[data-bs-theme="dark"] .commons-run-code .hl.sng { color: #98c379; } -.commons-run-r-details > .commons-run-r-code { +.commons-run-details > .commons-run-code { margin-top: 0.4rem; } -.commons-run-r-plot { +.commons-run-plot { border: 1px solid var(--bs-border-color, #dee2e6); border-radius: 0.5rem; height: auto;