diff --git a/.cursor/worktrees.json b/.cursor/worktrees.json new file mode 100644 index 00000000..e691ad38 --- /dev/null +++ b/.cursor/worktrees.json @@ -0,0 +1,6 @@ +{ + "setup-worktree": [ + "pdm install", + "cp $ROOT_WORKTREE_PATH/.env .env" + ] +} diff --git a/askui_chat.db b/askui_chat.db new file mode 100644 index 00000000..3ef4ec47 Binary files /dev/null and b/askui_chat.db differ diff --git a/debug_db.py b/debug_db.py new file mode 100644 index 00000000..322a5c90 --- /dev/null +++ b/debug_db.py @@ -0,0 +1,39 @@ +"""Debug test to understand the database issue.""" + +from sqlalchemy import create_engine, text + +from askui.chat.api.db.base import Base + +# Import all models + + +def test_debug_database(): + """Debug test to check database table creation.""" + engine = create_engine("sqlite:///:memory:") + + # Create tables + Base.metadata.create_all(engine) + + # Check if tables exist + with engine.connect() as conn: + result = conn.execute(text("SELECT name FROM sqlite_master WHERE type='table'")) + tables = [row[0] for row in result] + print(f"Tables created: {tables}") + + # Check if assistants table exists + if "assistants" in tables: + print("✅ assistants table exists") + else: + print("❌ assistants table missing") + + # Try to query assistants table + try: + result = conn.execute(text("SELECT COUNT(*) FROM assistants")) + count = result.scalar() + print(f"✅ Query successful, count: {count}") + except Exception as e: + print(f"❌ Query failed: {e}") + + +if __name__ == "__main__": + test_debug_database() diff --git a/pdm.lock b/pdm.lock index 00dfa448..f2482175 100644 --- a/pdm.lock +++ b/pdm.lock @@ -5,7 +5,7 @@ groups = ["default", "all", "android", "chat", "dev", "pynput", "web"] strategy = ["inherit_metadata"] lock_version = "4.5.0" -content_hash = "sha256:1d2d76cd5a60e5e1bd0b28e095e97c88afd9ac61ba00468b91168df6455b7494" +content_hash = "sha256:6f538558b17baad304306ceaf8793362d3eb516e583d31b9f0dbb4138b465f0f" [[metadata.targets]] requires_python = ">=3.10" @@ -21,6 +21,20 @@ files = [ {file = "aiofiles-24.1.0.tar.gz", hash = "sha256:22a075c9e5a3810f0c2e48f3008c94d68c65d763b9b03857924c99e57355166c"}, ] +[[package]] +name = "aiosqlite" +version = "0.21.0" +requires_python = ">=3.9" +summary = "asyncio bridge to the standard sqlite3 module" +groups = ["default"] +dependencies = [ + "typing-extensions>=4.0", +] +files = [ + {file = "aiosqlite-0.21.0-py3-none-any.whl", hash = "sha256:2549cf4057f95f53dcba16f2b64e8e2791d7e1adedb13197dd8ed77bb226d7d0"}, + {file = "aiosqlite-0.21.0.tar.gz", hash = "sha256:131bb8056daa3bc875608c631c678cda73922a2d4ba8aec373b19f18c17e7aa3"}, +] + [[package]] name = "annotated-types" version = "0.7.0" @@ -89,7 +103,7 @@ name = "asgi-correlation-id" version = "4.3.4" requires_python = "<4.0,>=3.8" summary = "Middleware correlating project logs to individual requests" -groups = ["all", "chat"] +groups = ["default", "all", "chat"] dependencies = [ "packaging", "starlette>=0.18", @@ -1039,7 +1053,7 @@ name = "greenlet" version = "3.2.4" requires_python = ">=3.9" summary = "Lightweight in-process concurrent programming" -groups = ["all", "chat", "dev", "web"] +groups = ["default", "all", "chat", "dev", "web"] files = [ {file = "greenlet-3.2.4-cp310-cp310-macosx_11_0_universal2.whl", hash = "sha256:8c68325b0d0acf8d91dde4e6f930967dd52a5302cd4062932a6b2e7c2969f47c"}, {file = "greenlet-3.2.4-cp310-cp310-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:94385f101946790ae13da500603491f04a76b6e4c059dab271b3ce2e283b2590"}, @@ -3416,6 +3430,54 @@ files = [ {file = "soupsieve-2.8.tar.gz", hash = "sha256:e2dd4a40a628cb5f28f6d4b0db8800b8f581b65bb380b97de22ba5ca8d72572f"}, ] +[[package]] +name = "sqlalchemy" +version = "2.0.43" +requires_python = ">=3.7" +summary = "Database Abstraction Library" +groups = ["default"] +dependencies = [ + "greenlet>=1; (platform_machine == \"win32\" or platform_machine == \"WIN32\" or platform_machine == \"AMD64\" or platform_machine == \"amd64\" or platform_machine == \"x86_64\" or platform_machine == \"ppc64le\" or platform_machine == \"aarch64\") and python_version < \"3.14\"", + "importlib-metadata; python_version < \"3.8\"", + "typing-extensions>=4.6.0", +] +files = [ + {file = "sqlalchemy-2.0.43-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:70322986c0c699dca241418fcf18e637a4369e0ec50540a2b907b184c8bca069"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:87accdbba88f33efa7b592dc2e8b2a9c2cdbca73db2f9d5c510790428c09c154"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:c00e7845d2f692ebfc7d5e4ec1a3fd87698e4337d09e58d6749a16aedfdf8612"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:022e436a1cb39b13756cf93b48ecce7aa95382b9cfacceb80a7d263129dfd019"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-musllinux_1_2_aarch64.whl", hash = "sha256:c5e73ba0d76eefc82ec0219d2301cb33bfe5205ed7a2602523111e2e56ccbd20"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-musllinux_1_2_x86_64.whl", hash = "sha256:9c2e02f06c68092b875d5cbe4824238ab93a7fa35d9c38052c033f7ca45daa18"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-win32.whl", hash = "sha256:e7a903b5b45b0d9fa03ac6a331e1c1d6b7e0ab41c63b6217b3d10357b83c8b00"}, + {file = "sqlalchemy-2.0.43-cp310-cp310-win_amd64.whl", hash = "sha256:4bf0edb24c128b7be0c61cd17eef432e4bef507013292415f3fb7023f02b7d4b"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-macosx_10_9_x86_64.whl", hash = "sha256:52d9b73b8fb3e9da34c2b31e6d99d60f5f99fd8c1225c9dad24aeb74a91e1d29"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:f42f23e152e4545157fa367b2435a1ace7571cab016ca26038867eb7df2c3631"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:4fb1a8c5438e0c5ea51afe9c6564f951525795cf432bed0c028c1cb081276685"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:db691fa174e8f7036afefe3061bc40ac2b770718be2862bfb03aabae09051aca"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-musllinux_1_2_aarch64.whl", hash = "sha256:fe2b3b4927d0bc03d02ad883f402d5de201dbc8894ac87d2e981e7d87430e60d"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-musllinux_1_2_x86_64.whl", hash = "sha256:4d3d9b904ad4a6b175a2de0738248822f5ac410f52c2fd389ada0b5262d6a1e3"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-win32.whl", hash = "sha256:5cda6b51faff2639296e276591808c1726c4a77929cfaa0f514f30a5f6156921"}, + {file = "sqlalchemy-2.0.43-cp311-cp311-win_amd64.whl", hash = "sha256:c5d1730b25d9a07727d20ad74bc1039bbbb0a6ca24e6769861c1aa5bf2c4c4a8"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:20d81fc2736509d7a2bd33292e489b056cbae543661bb7de7ce9f1c0cd6e7f24"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:25b9fc27650ff5a2c9d490c13c14906b918b0de1f8fcbb4c992712d8caf40e83"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:6772e3ca8a43a65a37c88e2f3e2adfd511b0b1da37ef11ed78dea16aeae85bd9"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:1a113da919c25f7f641ffbd07fbc9077abd4b3b75097c888ab818f962707eb48"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:4286a1139f14b7d70141c67a8ae1582fc2b69105f1b09d9573494eb4bb4b2687"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:529064085be2f4d8a6e5fab12d36ad44f1909a18848fcfbdb59cc6d4bbe48efe"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-win32.whl", hash = "sha256:b535d35dea8bbb8195e7e2b40059e2253acb2b7579b73c1b432a35363694641d"}, + {file = "sqlalchemy-2.0.43-cp312-cp312-win_amd64.whl", hash = "sha256:1c6d85327ca688dbae7e2b06d7d84cfe4f3fffa5b5f9e21bb6ce9d0e1a0e0e0a"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:e7c08f57f75a2bb62d7ee80a89686a5e5669f199235c6d1dac75cd59374091c3"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:14111d22c29efad445cd5021a70a8b42f7d9152d8ba7f73304c4d82460946aaa"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:21b27b56eb2f82653168cefe6cb8e970cdaf4f3a6cb2c5e3c3c1cf3158968ff9"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:9c5a9da957c56e43d72126a3f5845603da00e0293720b03bde0aacffcf2dc04f"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:5d79f9fdc9584ec83d1b3c75e9f4595c49017f5594fee1a2217117647225d738"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:9df7126fd9db49e3a5a3999442cc67e9ee8971f3cb9644250107d7296cb2a164"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-win32.whl", hash = "sha256:7f1ac7828857fcedb0361b48b9ac4821469f7694089d15550bbcf9ab22564a1d"}, + {file = "sqlalchemy-2.0.43-cp313-cp313-win_amd64.whl", hash = "sha256:971ba928fcde01869361f504fcff3b7143b47d30de188b11c6357c0505824197"}, + {file = "sqlalchemy-2.0.43-py3-none-any.whl", hash = "sha256:1681c21dd2ccee222c2fe0bef671d1aef7c504087c9c4e800371cfcc8ac966fc"}, + {file = "sqlalchemy-2.0.43.tar.gz", hash = "sha256:788bfcef6787a7764169cfe9859fe425bf44559619e1d9f56f5bddf2ebf6f417"}, +] + [[package]] name = "sse-starlette" version = "3.0.2" diff --git a/pyproject.toml b/pyproject.toml index 5bf44deb..c3657fc8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,6 +31,9 @@ dependencies = [ "bson>=0.5.10", "aiofiles>=24.1.0", "anyio==4.10.0", # We need to pin this version otherwise listing mcp tools using fastmcp within runner fails + "sqlalchemy>=2.0.0", + "aiosqlite>=0.19.0", + "asgi-correlation-id>=0.2.0", ] requires-python = ">=3.10" readme = "README.md" diff --git a/src/askui/chat/__main__.py b/src/askui/chat/__main__.py index 0e7d7772..0ad580d4 100644 --- a/src/askui/chat/__main__.py +++ b/src/askui/chat/__main__.py @@ -1,11 +1,19 @@ import uvicorn - from askui.chat.api.app import app from askui.chat.api.dependencies import get_settings from askui.chat.api.telemetry.integrations.fastapi import instrument +from askui.chat.migrations import MigrationRunner if __name__ == "__main__": settings = get_settings() + + # Run migration if needed + runner = MigrationRunner(settings.db.url) + if runner.should_migrate(settings.data_dir): + print("Starting database migration...") + runner.migrate(settings.data_dir) + print("Migration completed") + instrument(app, settings.telemetry) uvicorn.run( app, diff --git a/src/askui/chat/api/app.py b/src/askui/chat/api/app.py index d5fa2080..e47368db 100644 --- a/src/askui/chat/api/app.py +++ b/src/askui/chat/api/app.py @@ -1,20 +1,17 @@ from contextlib import asynccontextmanager from typing import AsyncGenerator -from fastapi import APIRouter, FastAPI, HTTPException, Request, status -from fastapi.middleware.cors import CORSMiddleware -from fastapi.responses import JSONResponse -from fastmcp import FastMCP - -from askui.chat.api.assistants.dependencies import get_assistant_service from askui.chat.api.assistants.router import router as assistants_router -from askui.chat.api.dependencies import SetEnvFromHeadersDep, get_settings +from askui.chat.api.assistants.service import AssistantService +from askui.chat.api.dependencies import (SetEnvFromHeadersDep, + get_session_factory, get_settings) from askui.chat.api.files.router import router as files_router from askui.chat.api.health.router import router as health_router -from askui.chat.api.mcp_clients.dependencies import get_mcp_client_manager_manager +from askui.chat.api.mcp_clients.dependencies import \ + get_mcp_client_manager_manager from askui.chat.api.mcp_clients.manager import McpServerConnectionError -from askui.chat.api.mcp_configs.dependencies import get_mcp_config_service from askui.chat.api.mcp_configs.router import router as mcp_configs_router +from askui.chat.api.mcp_configs.service import McpConfigService from askui.chat.api.mcp_servers.android import mcp as android_mcp from askui.chat.api.mcp_servers.computer import mcp as computer_mcp from askui.chat.api.mcp_servers.testing import mcp as testing_mcp @@ -23,23 +20,40 @@ from askui.chat.api.runs.router import router as runs_router from askui.chat.api.threads.router import router as threads_router from askui.chat.api.workflows.router import router as workflows_router -from askui.utils.api_utils import ( - ConflictError, - FileTooLargeError, - ForbiddenError, - LimitReachedError, - NotFoundError, -) +from askui.utils.api_utils import (ConflictError, FileTooLargeError, + ForbiddenError, LimitReachedError, + NotFoundError) +from fastapi import APIRouter, FastAPI, HTTPException, Request, status +from fastapi.middleware.cors import CORSMiddleware +from fastapi.responses import JSONResponse +from fastmcp import FastMCP settings = get_settings() @asynccontextmanager async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: # noqa: ARG001 - assistant_service = get_assistant_service(settings=settings) + # Get settings and create session factory + settings = get_settings() + session_factory = get_session_factory(settings) + + # Run migrations if needed + from askui.chat.migrations import MigrationRunner + runner = MigrationRunner(settings.db.url) + if runner.should_migrate(settings.data_dir): + print("Starting database migration...") + runner.migrate(settings.data_dir) + print("Migration completed") + + # Create services manually for seeding + assistant_service = AssistantService(session_factory) assistant_service.seed() - mcp_config_service = get_mcp_config_service(settings=settings) + + mcp_config_service = McpConfigService( + session_factory, settings.data_dir, settings.mcp_configs + ) mcp_config_service.seed() + yield await get_mcp_client_manager_manager(mcp_config_service).disconnect_all(force=True) @@ -174,3 +188,4 @@ def mcp_server_connection_error_handler( allow_methods=["*"], allow_headers=["*"], ) +) diff --git a/src/askui/chat/api/assistants/dependencies.py b/src/askui/chat/api/assistants/dependencies.py index d0d99dfb..21a5a287 100644 --- a/src/askui/chat/api/assistants/dependencies.py +++ b/src/askui/chat/api/assistants/dependencies.py @@ -1,13 +1,16 @@ -from fastapi import Depends +from typing import Callable from askui.chat.api.assistants.service import AssistantService -from askui.chat.api.dependencies import SettingsDep -from askui.chat.api.settings import Settings +from askui.chat.api.dependencies import SessionFactoryDep +from fastapi import Depends +from sqlalchemy.orm import Session -def get_assistant_service(settings: Settings = SettingsDep) -> AssistantService: +def get_assistant_service( + session_factory: Callable[[], Session] = SessionFactoryDep, +) -> AssistantService: """Get AssistantService instance.""" - return AssistantService(settings.data_dir) + return AssistantService(session_factory) AssistantServiceDep = Depends(get_assistant_service) diff --git a/src/askui/chat/api/assistants/models.py b/src/askui/chat/api/assistants/models.py index 9d3a23aa..ae768fc4 100644 --- a/src/askui/chat/api/assistants/models.py +++ b/src/askui/chat/api/assistants/models.py @@ -1,59 +1,78 @@ -from typing import Literal +"""Assistant database model.""" -from pydantic import BaseModel, Field +from datetime import datetime, timezone -from askui.chat.api.models import AssistantId, WorkspaceId, WorkspaceResource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id -from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven +from askui.chat.api.assistants.schemas import Assistant, AssistantCreateParams +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import AssistantId +from bson import ObjectId +from sqlalchemy import JSON, Column, DateTime, String, Text -class AssistantBase(BaseModel): - """Base assistant model.""" +class AssistantModel(Base): + """Assistant database model.""" - name: str | None = None - description: str | None = None - avatar: str | None = None - tools: list[str] = Field(default_factory=list) - system: str | None = None + __tablename__ = "assistants" + id = Column(AssistantId, primary_key=True) + workspace_id = Column(String(36), nullable=True, index=True) + created_at = Column(DateTime, nullable=False, index=True) + name = Column(String, nullable=True) + description = Column(String, nullable=True) + avatar = Column(Text, nullable=True) + tools = Column(JSON, nullable=False) + system = Column(Text, nullable=True) + @staticmethod + def create_id() -> str: + """Create a new assistant ID with prefix.""" + return f"asst_{ObjectId()}" -class AssistantCreateParams(AssistantBase): - """Parameters for creating an assistant.""" + def to_pydantic(self) -> Assistant: + """Convert to Pydantic model.""" + # Ensure created_at is timezone-aware + created_at = self.created_at + if created_at.tzinfo is None: + created_at = created_at.replace(tzinfo=timezone.utc) - -class AssistantModifyParams(BaseModelWithNotGiven): - """Parameters for modifying an assistant.""" - - name: str | NotGiven = NOT_GIVEN - description: str | NotGiven = NOT_GIVEN - avatar: str | NotGiven = NOT_GIVEN - tools: list[str] | NotGiven = NOT_GIVEN - system: str | NotGiven = NOT_GIVEN - - -class Assistant(AssistantBase, WorkspaceResource): - """An assistant that can be used in a thread.""" - - id: AssistantId - object: Literal["assistant"] = "assistant" - created_at: UnixDatetime + return Assistant( + id=self.id, # Prefix is handled by the specialized type + workspace_id=self.workspace_id, + created_at=created_at, + name=self.name, + description=self.description, + avatar=self.avatar, + tools=self.tools, + system=self.system, + ) @classmethod - def create( - cls, workspace_id: WorkspaceId, params: AssistantCreateParams - ) -> "Assistant": + def from_pydantic(cls, assistant: Assistant) -> "AssistantModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("asst"), - created_at=now(), - workspace_id=workspace_id, - **params.model_dump(), + id=assistant.id, + workspace_id=str(assistant.workspace_id) + if assistant.workspace_id + else None, + created_at=assistant.created_at, + name=assistant.name, + description=assistant.description, + avatar=assistant.avatar, + tools=assistant.tools, + system=assistant.system, ) - def modify(self, params: AssistantModifyParams) -> "Assistant": - return Assistant.model_validate( - { - **self.model_dump(), - **params.model_dump(), - } + @classmethod + def from_create_params( + cls, params: AssistantCreateParams, workspace_id: str | None = None + ) -> "AssistantModel": + """Create from create parameters.""" + return cls( + id=cls.create_id(), + workspace_id=str(workspace_id) if workspace_id else None, + created_at=datetime.now(timezone.utc), + name=params.name, + description=params.description, + avatar=params.avatar, + tools=params.tools, + system=params.system, ) diff --git a/src/askui/chat/api/assistants/router.py b/src/askui/chat/api/assistants/router.py index 76ae8ab4..04e3e397 100644 --- a/src/askui/chat/api/assistants/router.py +++ b/src/askui/chat/api/assistants/router.py @@ -1,9 +1,7 @@ from typing import Annotated -from fastapi import APIRouter, Header, status - from askui.chat.api.assistants.dependencies import AssistantServiceDep -from askui.chat.api.assistants.models import ( +from askui.chat.api.assistants.schemas import ( Assistant, AssistantCreateParams, AssistantModifyParams, @@ -12,6 +10,7 @@ from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.models import AssistantId, WorkspaceId from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, Header, status router = APIRouter(prefix="/assistants", tags=["assistants"]) diff --git a/src/askui/chat/api/assistants/schemas.py b/src/askui/chat/api/assistants/schemas.py new file mode 100644 index 00000000..9d3a23aa --- /dev/null +++ b/src/askui/chat/api/assistants/schemas.py @@ -0,0 +1,59 @@ +from typing import Literal + +from pydantic import BaseModel, Field + +from askui.chat.api.models import AssistantId, WorkspaceId, WorkspaceResource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id +from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven + + +class AssistantBase(BaseModel): + """Base assistant model.""" + + name: str | None = None + description: str | None = None + avatar: str | None = None + tools: list[str] = Field(default_factory=list) + system: str | None = None + + +class AssistantCreateParams(AssistantBase): + """Parameters for creating an assistant.""" + + +class AssistantModifyParams(BaseModelWithNotGiven): + """Parameters for modifying an assistant.""" + + name: str | NotGiven = NOT_GIVEN + description: str | NotGiven = NOT_GIVEN + avatar: str | NotGiven = NOT_GIVEN + tools: list[str] | NotGiven = NOT_GIVEN + system: str | NotGiven = NOT_GIVEN + + +class Assistant(AssistantBase, WorkspaceResource): + """An assistant that can be used in a thread.""" + + id: AssistantId + object: Literal["assistant"] = "assistant" + created_at: UnixDatetime + + @classmethod + def create( + cls, workspace_id: WorkspaceId, params: AssistantCreateParams + ) -> "Assistant": + return cls( + id=generate_time_ordered_id("asst"), + created_at=now(), + workspace_id=workspace_id, + **params.model_dump(), + ) + + def modify(self, params: AssistantModifyParams) -> "Assistant": + return Assistant.model_validate( + { + **self.model_dump(), + **params.model_dump(), + } + ) diff --git a/src/askui/chat/api/assistants/service.py b/src/askui/chat/api/assistants/service.py index 3c4248fc..2e78f426 100644 --- a/src/askui/chat/api/assistants/service.py +++ b/src/askui/chat/api/assistants/service.py @@ -1,81 +1,84 @@ -from pathlib import Path +from typing import Callable -from askui.chat.api.assistants.models import ( +from askui.chat.api.assistants.models import AssistantModel +from askui.chat.api.assistants.schemas import ( Assistant, AssistantCreateParams, AssistantModifyParams, ) from askui.chat.api.assistants.seeds import SEEDS +from askui.chat.api.db.query_builder import QueryBuilder from askui.chat.api.models import AssistantId, WorkspaceId -from askui.chat.api.utils import build_workspace_filter_fn -from askui.utils.api_utils import ( - LIST_LIMIT_MAX, - ConflictError, - ForbiddenError, - ListQuery, - ListResponse, - NotFoundError, - list_resources, -) +from askui.utils.api_utils import ForbiddenError, ListQuery, ListResponse, NotFoundError +from askui.utils.not_given import NOT_GIVEN +from sqlalchemy.orm import Session class AssistantService: - def __init__(self, base_dir: Path) -> None: - self._base_dir = base_dir - self._assistants_dir = base_dir / "assistants" - - def _get_assistant_path(self, assistant_id: AssistantId, new: bool = False) -> Path: - assistant_path = self._assistants_dir / f"{assistant_id}.json" - exists = assistant_path.exists() - if new and exists: - error_msg = f"Assistant {assistant_id} already exists" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"Assistant {assistant_id} not found" - raise NotFoundError(error_msg) - return assistant_path + def __init__(self, session_factory: Callable[[], Session]) -> None: + self._session_factory = session_factory + + def _to_pydantic(self, db_model: AssistantModel) -> Assistant: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_( self, workspace_id: WorkspaceId | None, query: ListQuery ) -> ListResponse[Assistant]: - return list_resources( - self._assistants_dir, - query, - Assistant, - filter_fn=build_workspace_filter_fn(workspace_id, Assistant), - ) + with self._session_factory() as session: + q = session.query(AssistantModel) + + # Filter by workspace + if workspace_id is not None: + q = q.filter(AssistantModel.workspace_id == str(workspace_id)) + else: + q = q.filter(AssistantModel.workspace_id.is_(None)) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, AssistantModel, query, AssistantModel.created_at, AssistantModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve( self, workspace_id: WorkspaceId | None, assistant_id: AssistantId ) -> Assistant: - try: - assistant_path = self._get_assistant_path(assistant_id) - content = assistant_path.read_text() - if not content.strip(): + with self._session_factory() as session: + db_assistant = ( + session.query(AssistantModel) + .filter(AssistantModel.id == assistant_id) + .first() + ) + if not db_assistant: error_msg = f"Assistant {assistant_id} not found" raise NotFoundError(error_msg) - assistant = Assistant.model_validate_json(content) + + # Check workspace access if not ( - assistant.workspace_id is None or assistant.workspace_id == workspace_id + db_assistant.workspace_id is None + or db_assistant.workspace_id == str(workspace_id) ): error_msg = f"Assistant {assistant_id} not found" raise NotFoundError(error_msg) - except FileNotFoundError as e: - error_msg = f"Assistant {assistant_id} not found" - raise NotFoundError(error_msg) from e - except (ValueError, TypeError) as e: - # Handle JSON parsing errors - error_msg = f"Assistant {assistant_id} not found" - raise NotFoundError(error_msg) from e - else: - return assistant + + return self._to_pydantic(db_assistant) def create( self, workspace_id: WorkspaceId, params: AssistantCreateParams ) -> Assistant: - assistant = Assistant.create(workspace_id, params) - self._save(assistant, new=True) - return assistant + with self._session_factory() as session: + db_assistant = AssistantModel.from_create_params(params, workspace_id) + session.add(db_assistant) + session.commit() + session.refresh(db_assistant) + + return self._to_pydantic(db_assistant) def modify( self, @@ -83,13 +86,44 @@ def modify( assistant_id: AssistantId, params: AssistantModifyParams, ) -> Assistant: - assistant = self.retrieve(workspace_id, assistant_id) - if assistant.workspace_id is None: - error_msg = f"Default assistant {assistant_id} cannot be modified" - raise ForbiddenError(error_msg) - modified = assistant.modify(params) - self._save(modified) - return modified + with self._session_factory() as session: + db_assistant = ( + session.query(AssistantModel) + .filter(AssistantModel.id == assistant_id) + .first() + ) + if not db_assistant: + error_msg = f"Assistant {assistant_id} not found" + raise NotFoundError(error_msg) + + # Check workspace access + if not ( + db_assistant.workspace_id is None + or db_assistant.workspace_id == str(workspace_id) + ): + error_msg = f"Assistant {assistant_id} not found" + raise NotFoundError(error_msg) + + if db_assistant.workspace_id is None: + error_msg = f"Default assistant {assistant_id} cannot be modified" + raise ForbiddenError(error_msg) + + # Update fields + if params.name is not NOT_GIVEN: + db_assistant.name = params.name + if params.description is not NOT_GIVEN: + db_assistant.description = params.description + if params.avatar is not NOT_GIVEN: + db_assistant.avatar = params.avatar + if params.tools is not NOT_GIVEN: + db_assistant.tools = params.tools + if params.system is not NOT_GIVEN: + db_assistant.system = params.system + + session.commit() + session.refresh(db_assistant) + + return self._to_pydantic(db_assistant) def delete( self, @@ -97,35 +131,67 @@ def delete( assistant_id: AssistantId, force: bool = False, ) -> None: - try: - assistant = self.retrieve(workspace_id, assistant_id) - if assistant.workspace_id is None and not force: + with self._session_factory() as session: + db_assistant = ( + session.query(AssistantModel) + .filter(AssistantModel.id == assistant_id) + .first() + ) + + if not db_assistant: + if not force: + error_msg = f"Assistant {assistant_id} not found" + raise NotFoundError(error_msg) + return + + # Check workspace access + if not ( + db_assistant.workspace_id is None + or db_assistant.workspace_id == str(workspace_id) + ): + if not force: + error_msg = f"Assistant {assistant_id} not found" + raise NotFoundError(error_msg) + return + + if db_assistant.workspace_id is None and not force: error_msg = f"Default assistant {assistant_id} cannot be deleted" raise ForbiddenError(error_msg) - try: - self._get_assistant_path(assistant_id).unlink() - except FileNotFoundError: - # File already deleted, that's fine - pass - except FileNotFoundError as e: - error_msg = f"Assistant {assistant_id} not found" - raise NotFoundError(error_msg) from e - except NotFoundError: - # If force=True and assistant doesn't exist, just ignore - if not force: - raise - # For force=True, we can ignore the NotFoundError - - def _save(self, assistant: Assistant, new: bool = False) -> None: - self._assistants_dir.mkdir(parents=True, exist_ok=True) - assistant_file = self._get_assistant_path(assistant.id, new=new) - assistant_file.write_text(assistant.model_dump_json(), encoding="utf-8") + + session.delete(db_assistant) + session.commit() def seed(self) -> None: """Seed the assistant service with default assistants.""" - for seed in SEEDS: - self.delete(None, seed.id, force=True) - try: - self._save(seed, new=True) - except ConflictError: # noqa: PERF203 - self._save(seed) + with self._session_factory() as session: + for seed in SEEDS: + # Check if already exists + existing = ( + session.query(AssistantModel) + .filter(AssistantModel.id == seed.id) + .first() + ) + + if existing: + # Update existing + existing.name = seed.name + existing.description = seed.description + existing.avatar = seed.avatar + existing.tools = seed.tools + existing.system = seed.system + else: + # Create new + db_assistant = AssistantModel( + id=seed.id, + workspace_id=None, # Default assistants have no workspace + created_at=seed.created_at, + name=seed.name, + description=seed.description, + avatar=seed.avatar, + tools=seed.tools, + system=seed.system, + ) + session.add(db_assistant) + + session.commit() + session.commit() diff --git a/src/askui/chat/api/db/__init__.py b/src/askui/chat/api/db/__init__.py new file mode 100644 index 00000000..30cc75fc --- /dev/null +++ b/src/askui/chat/api/db/__init__.py @@ -0,0 +1,29 @@ +"""Database module for chat API.""" + +from .base import Base +from .session import get_session_factory +from .types import ( + AssistantId, + FileId, + McpConfigId, + MessageId, + PrefixedObjectId, + RunId, + ThreadId, + WorkflowId, + create_prefixed_id_type, +) + +__all__ = [ + "Base", + "get_session_factory", + "AssistantId", + "FileId", + "McpConfigId", + "MessageId", + "PrefixedObjectId", + "RunId", + "ThreadId", + "WorkflowId", + "create_prefixed_id_type", +] diff --git a/src/askui/chat/api/db/base.py b/src/askui/chat/api/db/base.py new file mode 100644 index 00000000..6869fed2 --- /dev/null +++ b/src/askui/chat/api/db/base.py @@ -0,0 +1,5 @@ +"""SQLAlchemy declarative base for chat API.""" + +from sqlalchemy.orm import declarative_base + +Base = declarative_base() diff --git a/src/askui/chat/api/db/engine.py b/src/askui/chat/api/db/engine.py new file mode 100644 index 00000000..b690cbce --- /dev/null +++ b/src/askui/chat/api/db/engine.py @@ -0,0 +1,18 @@ +"""Database engine configuration for chat API.""" + +from pathlib import Path + +from sqlalchemy import create_engine +from sqlalchemy.engine import Engine + + +def create_database_engine(db_path: Path) -> Engine: + """Create SQLAlchemy engine for SQLite database. + + Args: + db_path (Path): Path to SQLite database file. + + Returns: + Engine: Configured SQLAlchemy engine. + """ + return create_engine(f"sqlite:///{db_path}") diff --git a/src/askui/chat/api/db/models.py b/src/askui/chat/api/db/models.py new file mode 100644 index 00000000..6382bf88 --- /dev/null +++ b/src/askui/chat/api/db/models.py @@ -0,0 +1,10 @@ +"""SQLAlchemy ORM models for chat API - centralized imports.""" + +# This file is kept for backward compatibility but individual models +# should be imported directly from their respective modules to avoid +# circular imports. + +# Example imports: +# from askui.chat.api.assistants.models import AssistantModel +# from askui.chat.api.threads.models import ThreadModel +# etc. diff --git a/src/askui/chat/api/db/query_builder.py b/src/askui/chat/api/db/query_builder.py new file mode 100644 index 00000000..156d4cbf --- /dev/null +++ b/src/askui/chat/api/db/query_builder.py @@ -0,0 +1,98 @@ +"""Shared query building utilities for database operations.""" + +from typing import Any, Callable, TypeVar + +from askui.utils.api_utils import ListQuery, ListResponse +from sqlalchemy import Column, desc +from sqlalchemy.orm import Query + +ModelT = TypeVar("ModelT") + + +class QueryBuilder: + """Builder for common database queries.""" + + @staticmethod + def apply_list_query( + query: Query[ModelT], + model_class: type[ModelT], + list_query: ListQuery, + created_at_column: Column[Any], + id_column: Column[str] | None = None, + ) -> Query[ModelT]: + """Apply list query parameters to a SQLAlchemy query. + + Args: + query (Query[ModelT]): The base query to modify. + model_class (type[ModelT]): The model class for type hints. + list_query (ListQuery): The list query parameters. + created_at_column (Column[Any]): The created_at column for ordering. + id_column (Column[str] | None): The ID column for pagination. + + Returns: + Query[ModelT]: The modified query. + """ + # Apply ordering using created_at + if list_query.order == "desc": + query = query.order_by(desc(created_at_column)) + else: + query = query.order_by(created_at_column) + + # Apply pagination using created_at by looking up the created_at values + # for the given IDs + if list_query.after and id_column is not None: + # Look up the created_at value for the after ID + after_subquery = ( + query.session.query(created_at_column) + .filter(id_column == list_query.after) + .scalar_subquery() + ) + + if list_query.order == "desc": + query = query.filter(created_at_column < after_subquery) + else: + query = query.filter(created_at_column > after_subquery) + + if list_query.before and id_column is not None: + # Look up the created_at value for the before ID + before_subquery = ( + query.session.query(created_at_column) + .filter(id_column == list_query.before) + .scalar_subquery() + ) + + if list_query.order == "desc": + query = query.filter(created_at_column > before_subquery) + else: + query = query.filter(created_at_column < before_subquery) + + return query + + @staticmethod + def build_list_response( + results: list[ModelT], + limit: int | None, + to_pydantic_func: Callable[[ModelT], Any], + ) -> ListResponse[Any]: + """Build a ListResponse from query results. + + Args: + results (list[ModelT]): The query results. + limit (int | None): The limit that was applied. + to_pydantic_func (Callable[[ModelT], Any]): Function to convert model to Pydantic. + + Returns: + ListResponse[Any]: The list response. + """ + has_more = len(results) > (limit or 20) + if has_more: + results = results[: limit or 20] + + data = [to_pydantic_func(result) for result in results] + + return ListResponse( + data=data, + has_more=has_more, + first_id=data[0].id if data else None, + last_id=data[-1].id if data else None, + ) diff --git a/src/askui/chat/api/db/session.py b/src/askui/chat/api/db/session.py new file mode 100644 index 00000000..065a7104 --- /dev/null +++ b/src/askui/chat/api/db/session.py @@ -0,0 +1,41 @@ +"""Session management for chat API database.""" + +import logging +from contextlib import contextmanager +from typing import Callable, Generator + +from askui.chat.api.settings import Settings +from sqlalchemy import create_engine +from sqlalchemy.orm import Session, sessionmaker + +logger = logging.getLogger(__name__) + + +def get_session_factory( + settings: Settings, +) -> Callable[[], Generator[Session, None, None]]: + """Get SQLAlchemy session factory. + + Args: + settings (Settings): Application settings containing database URL. + + Returns: + Callable[[], Generator[Session, None, None]]: Session factory that returns a context manager. + """ + # Enable SQL logging if debug level is set + echo = logger.isEnabledFor(logging.DEBUG) + engine = create_engine(settings.db.url, echo=echo) + SessionLocal = sessionmaker(bind=engine, expire_on_commit=False) + + @contextmanager + def session_factory() -> Generator[Session, None, None]: + session = SessionLocal() + try: + yield session + except Exception: + session.rollback() + raise + finally: + session.close() + + return session_factory diff --git a/src/askui/chat/api/db/types.py b/src/askui/chat/api/db/types.py new file mode 100644 index 00000000..40bda7f1 --- /dev/null +++ b/src/askui/chat/api/db/types.py @@ -0,0 +1,65 @@ +"""Custom SQLAlchemy types for chat API.""" + +from typing import Any + +from sqlalchemy import String, TypeDecorator + + +class PrefixedObjectId(TypeDecorator): + """Custom type for storing BSON ObjectIds with prefixes in SQLite. + + Stores ObjectIds without prefix in the database. The service layer + is responsible for adding/removing prefixes when converting to/from + Pydantic models. + """ + + impl = String(24) + cache_ok = True + + def process_bind_param(self, value: Any, dialect: Any) -> str | None: + """Process value before storing in database.""" + if value is None: + return value + # Remove prefix before storing + if isinstance(value, str) and "_" in value: + return value.split("_", 1)[1] + return str(value) + + def process_result_value(self, value: str | None, dialect: Any) -> str | None: + """Process value when reading from database.""" + # Service layer will add prefix when converting to Pydantic + return value + + +def create_prefixed_id_type(prefix: str) -> type[PrefixedObjectId]: + """Create a specialized ObjectId type for a specific prefix. + + Args: + prefix (str): The prefix to use (e.g., "asst", "thread"). + + Returns: + type[PrefixedObjectId]: A specialized type class. + """ + + class SpecializedPrefixedObjectId(PrefixedObjectId): + """Specialized ObjectId type with prefix awareness.""" + + cache_ok = True + + def process_result_value(self, value: str | None, dialect: Any) -> str | None: + """Add prefix when reading from database.""" + if value is None: + return value + return f"{prefix}_{value}" + + return SpecializedPrefixedObjectId + + +# Specialized types for each resource +AssistantId = create_prefixed_id_type("asst") +ThreadId = create_prefixed_id_type("thread") +MessageId = create_prefixed_id_type("msg") +RunId = create_prefixed_id_type("run") +FileId = create_prefixed_id_type("file") +WorkflowId = create_prefixed_id_type("workflow") +McpConfigId = create_prefixed_id_type("mcp") diff --git a/src/askui/chat/api/dependencies.py b/src/askui/chat/api/dependencies.py index 7264f598..344fba0c 100644 --- a/src/askui/chat/api/dependencies.py +++ b/src/askui/chat/api/dependencies.py @@ -1,15 +1,13 @@ import os -from pathlib import Path from typing import Annotated, Optional +from askui.chat.api.db.session import get_session_factory +from askui.chat.api.settings import Settings +from askui.utils.api_utils import ListQuery from fastapi import Depends, Header from fastapi.security import APIKeyHeader, HTTPAuthorizationCredentials, HTTPBearer from pydantic import UUID4 -from askui.chat.api.models import WorkspaceId -from askui.chat.api.settings import Settings -from askui.utils.api_utils import ListQuery - def get_settings() -> Settings: """Get ChatApiSettings instance.""" @@ -60,14 +58,13 @@ def set_env_from_headers( SetEnvFromHeadersDep = Depends(set_env_from_headers) -def get_workspace_dir( - askui_workspace: Annotated[WorkspaceId, Header()], - settings: Settings = SettingsDep, -) -> Path: - return settings.data_dir / "workspaces" / str(askui_workspace) +ListQueryDep = Depends(ListQuery) -WorkspaceDirDep = Depends(get_workspace_dir) +def get_session_factory_dep(settings: Settings = SettingsDep): + """Get SQLAlchemy session factory dependency.""" + return get_session_factory(settings) -ListQueryDep = Depends(ListQuery) +SessionFactoryDep = Depends(get_session_factory_dep) +SessionFactoryDep = Depends(get_session_factory_dep) diff --git a/src/askui/chat/api/files/dependencies.py b/src/askui/chat/api/files/dependencies.py index 75f2f39c..19602265 100644 --- a/src/askui/chat/api/files/dependencies.py +++ b/src/askui/chat/api/files/dependencies.py @@ -1,14 +1,11 @@ -from pathlib import Path - -from fastapi import Depends - -from askui.chat.api.dependencies import WorkspaceDirDep +from askui.chat.api.dependencies import SessionFactoryDep from askui.chat.api.files.service import FileService +from fastapi import Depends -def get_file_service(workspace_dir: Path = WorkspaceDirDep) -> FileService: +def get_file_service(session_factory=SessionFactoryDep) -> FileService: """Get FileService instance.""" - return FileService(workspace_dir) + return FileService(session_factory) FileServiceDep = Depends(get_file_service) diff --git a/src/askui/chat/api/files/models.py b/src/askui/chat/api/files/models.py index cf55c127..256dbb6c 100644 --- a/src/askui/chat/api/files/models.py +++ b/src/askui/chat/api/files/models.py @@ -1,42 +1,59 @@ -import mimetypes -from typing import Literal - -from pydantic import BaseModel, Field - -from askui.chat.api.models import FileId -from askui.utils.api_utils import Resource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id - - -class FileBase(BaseModel): - """Base file model.""" - - size: int = Field(description="In bytes", ge=0) - media_type: str - - -class FileCreateParams(FileBase): - filename: str | None = None - - -class File(FileBase, Resource): - """A file that can be stored and managed.""" - - id: FileId - object: Literal["file"] = "file" - created_at: UnixDatetime - filename: str = Field(min_length=1) +"""File database model.""" + +from datetime import datetime, timezone + +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import FileId +from askui.chat.api.files.schemas import File as FileSchema +from bson import ObjectId +from sqlalchemy import Column, DateTime, Integer, String + + +class FileModel(Base): + """File database model.""" + + __tablename__ = "files" + id = Column(FileId, primary_key=True) + created_at = Column(DateTime, nullable=False, index=True) + filename = Column(String, nullable=False) + size = Column(Integer, nullable=False) + media_type = Column(String, nullable=False) + + @staticmethod + def create_id() -> str: + """Create a new file ID with prefix.""" + return f"file_{ObjectId()}" + + def to_pydantic(self) -> FileSchema: + """Convert to Pydantic model.""" + return FileSchema( + id=self.id, # Prefix is handled by the specialized type + created_at=self.created_at, + filename=self.filename, + size=self.size, + media_type=self.media_type, + ) @classmethod - def create(cls, params: FileCreateParams) -> "File": - id_ = generate_time_ordered_id("file") - filename = ( - params.filename or f"{id_}{mimetypes.guess_extension(params.media_type)}" + def from_pydantic(cls, file: FileSchema) -> "FileModel": + """Create from Pydantic model.""" + return cls( + id=file.id, + created_at=file.created_at, + filename=file.filename, + size=file.size, + media_type=file.media_type, ) + + @classmethod + def from_create_params( + cls, filename: str, size: int, media_type: str + ) -> "FileModel": + """Create from create parameters.""" return cls( - id=id_, - created_at=now(), + id=cls.create_id(), + created_at=datetime.now(timezone.utc), filename=filename, - **params.model_dump(exclude={"filename"}), + size=size, + media_type=media_type, ) diff --git a/src/askui/chat/api/files/router.py b/src/askui/chat/api/files/router.py index 3ebcf8cd..93066730 100644 --- a/src/askui/chat/api/files/router.py +++ b/src/askui/chat/api/files/router.py @@ -1,12 +1,11 @@ -from fastapi import APIRouter, UploadFile, status -from fastapi.responses import FileResponse - from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.files.dependencies import FileServiceDep -from askui.chat.api.files.models import File as FileModel +from askui.chat.api.files.schemas import File as FileSchema from askui.chat.api.files.service import FileService from askui.chat.api.models import FileId from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, UploadFile, status +from fastapi.responses import FileResponse router = APIRouter(prefix="/files", tags=["files"]) @@ -15,7 +14,7 @@ def list_files( query: ListQuery = ListQueryDep, file_service: FileService = FileServiceDep, -) -> ListResponse[FileModel]: +) -> ListResponse[FileSchema]: """List all files.""" return file_service.list_(query=query) @@ -24,7 +23,7 @@ def list_files( async def upload_file( file: UploadFile, file_service: FileService = FileServiceDep, -) -> FileModel: +) -> FileSchema: """Upload a new file.""" return await file_service.upload_file(file) @@ -33,7 +32,7 @@ async def upload_file( def retrieve_file( file_id: FileId, file_service: FileService = FileServiceDep, -) -> FileModel: +) -> FileSchema: """Get file metadata by ID.""" return file_service.retrieve(file_id) diff --git a/src/askui/chat/api/files/schemas.py b/src/askui/chat/api/files/schemas.py new file mode 100644 index 00000000..cf55c127 --- /dev/null +++ b/src/askui/chat/api/files/schemas.py @@ -0,0 +1,42 @@ +import mimetypes +from typing import Literal + +from pydantic import BaseModel, Field + +from askui.chat.api.models import FileId +from askui.utils.api_utils import Resource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id + + +class FileBase(BaseModel): + """Base file model.""" + + size: int = Field(description="In bytes", ge=0) + media_type: str + + +class FileCreateParams(FileBase): + filename: str | None = None + + +class File(FileBase, Resource): + """A file that can be stored and managed.""" + + id: FileId + object: Literal["file"] = "file" + created_at: UnixDatetime + filename: str = Field(min_length=1) + + @classmethod + def create(cls, params: FileCreateParams) -> "File": + id_ = generate_time_ordered_id("file") + filename = ( + params.filename or f"{id_}{mimetypes.guess_extension(params.media_type)}" + ) + return cls( + id=id_, + created_at=now(), + filename=filename, + **params.model_dump(exclude={"filename"}), + ) diff --git a/src/askui/chat/api/files/service.py b/src/askui/chat/api/files/service.py index ee5aed9a..f43d1a71 100644 --- a/src/askui/chat/api/files/service.py +++ b/src/askui/chat/api/files/service.py @@ -1,151 +1,125 @@ +"""File service with SQLAlchemy persistence.""" + import logging -import mimetypes import shutil -import tempfile from pathlib import Path +from typing import Callable -from fastapi import UploadFile - -from askui.chat.api.files.models import File, FileCreateParams +from askui.chat.api.db.query_builder import QueryBuilder +from askui.chat.api.files.models import FileModel +from askui.chat.api.files.schemas import File from askui.chat.api.models import FileId from askui.utils.api_utils import ( - ConflictError, FileTooLargeError, ListQuery, ListResponse, NotFoundError, - list_resources, ) +from fastapi import UploadFile +from sqlalchemy.orm import Session logger = logging.getLogger(__name__) # Constants MAX_FILE_SIZE = 20 * 1024 * 1024 # 20MB supported -CHUNK_SIZE = 1024 * 1024 # 1MB for uploading and downloading class FileService: - """Service for managing File resources with filesystem persistence.""" - - def __init__(self, base_dir: Path) -> None: - self._base_dir = base_dir - self._files_dir = base_dir / "files" - self._static_dir = base_dir / "static" - - def _get_file_path(self, file_id: FileId, new: bool = False) -> Path: - """Get the path for file metadata.""" - file_path = self._files_dir / f"{file_id}.json" - exists = file_path.exists() - if new and exists: - error_msg = f"File {file_id} already exists" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"File {file_id} not found" - raise NotFoundError(error_msg) - return file_path - - def _get_static_file_path(self, file: File) -> Path: - """Get the path for the static file based on extension.""" - # For application/octet-stream, don't add .bin extension - extension = "" - if file.media_type != "application/octet-stream": - extension = mimetypes.guess_extension(file.media_type) or "" - return self._static_dir / f"{file.id}{extension}" + """Service for managing File resources with SQLAlchemy persistence.""" + + def __init__(self, session_factory: Callable[[], Session]) -> None: + self._session_factory = session_factory + self._files_dir = Path.cwd() / "chat" / "files" + self._static_dir = Path.cwd() / "chat" / "static" + + def _to_pydantic(self, db_model: FileModel) -> File: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_(self, query: ListQuery) -> ListResponse[File]: - """List files with pagination and filtering.""" - return list_resources(self._files_dir, query, File) + """List files with pagination.""" + with self._session_factory() as session: + q = session.query(FileModel) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, FileModel, query, FileModel.created_at, FileModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve(self, file_id: FileId) -> File: - """Retrieve file metadata by ID.""" - try: - file_path = self._get_file_path(file_id) - return File.model_validate_json(file_path.read_text()) - except FileNotFoundError as e: - error_msg = f"File {file_id} not found" - raise NotFoundError(error_msg) from e + """Retrieve a file by ID.""" + with self._session_factory() as session: + db_file = session.query(FileModel).filter(FileModel.id == file_id).first() + if not db_file: + error_msg = f"File {file_id} not found" + raise NotFoundError(error_msg) + return self._to_pydantic(db_file) + + def create( + self, filename: str, size: int, media_type: str, file: UploadFile + ) -> File: + """Create a new file.""" + # Check file size + if file.size and file.size > MAX_FILE_SIZE: + error_msg = ( + f"File size {file.size} exceeds maximum allowed size {MAX_FILE_SIZE}" + ) + raise FileTooLargeError(MAX_FILE_SIZE) + + with self._session_factory() as session: + # Create database record + db_file = FileModel.from_create_params(filename, size, media_type) + session.add(db_file) + session.commit() + session.refresh(db_file) + + # Save file content to filesystem + self._files_dir.mkdir(parents=True, exist_ok=True) + file_path = self._files_dir / db_file.id + with open(file_path, "wb") as f: + shutil.copyfileobj(file.file, f) + + return self._to_pydantic(db_file) def delete(self, file_id: FileId) -> None: - """Delete a file and its content. - - *Important*: We may be left with a static file that is not associated with any - file metadata if this fails. - """ - try: - file = self.retrieve(file_id) - static_path = self._get_static_file_path(file) - self._get_file_path(file_id).unlink() - if static_path.exists(): - static_path.unlink() - except FileNotFoundError as e: - error_msg = f"File {file_id} not found" - raise NotFoundError(error_msg) from e + """Delete a file.""" + with self._session_factory() as session: + db_file = session.query(FileModel).filter(FileModel.id == file_id).first() + if not db_file: + error_msg = f"File {file_id} not found" + raise NotFoundError(error_msg) + + # Delete file from filesystem + file_path = self._files_dir / file_id + if file_path.exists(): + file_path.unlink() + + # Delete database record + session.delete(db_file) + session.commit() + + def get_file_path(self, file_id: FileId) -> Path: + """Get the filesystem path for a file.""" + return self._files_dir / file_id + + async def upload_file(self, file: UploadFile) -> File: + """Upload a file (async wrapper for create).""" + filename = file.filename or "unknown" + size = file.size or 0 + media_type = file.content_type or "application/octet-stream" + return self.create(filename, size, media_type, file) def retrieve_file_content(self, file_id: FileId) -> tuple[File, Path]: - """Get file metadata and path for downloading.""" - file = self.retrieve(file_id) - static_path = self._get_static_file_path(file) - return file, static_path - - async def _write_to_temp_file( - self, - file: UploadFile, - ) -> tuple[FileCreateParams, Path]: - size = 0 - self._static_dir.mkdir(parents=True, exist_ok=True) - temp_file = tempfile.NamedTemporaryFile( - delete=False, - dir=self._static_dir, - suffix=".temp", - ) - temp_path = Path(temp_file.name) - with temp_file: - while chunk := await file.read(CHUNK_SIZE): - temp_file.write(chunk) - size += len(chunk) - if size > MAX_FILE_SIZE: - raise FileTooLargeError(MAX_FILE_SIZE) - mime_type = file.content_type or "application/octet-stream" - params = FileCreateParams( - filename=file.filename, - size=size, - media_type=mime_type, - ) - return params, temp_path - - def create(self, params: FileCreateParams, path: Path) -> File: - file_model = File.create(params) - self._static_dir.mkdir(parents=True, exist_ok=True) - static_path = self._get_static_file_path(file_model) - shutil.move(path, static_path) - self._save(file_model, new=True) - - return file_model - - async def upload_file( - self, - file: UploadFile, - ) -> File: - """Upload a file. - - *Important*: We may be left with a static file that is not associated with any - file metadata if this fails. - """ - temp_path: Path | None = None - try: - params, temp_path = await self._write_to_temp_file(file) - file_model = self.create(params, temp_path) - except Exception: - logger.exception("Failed to upload file") - raise - else: - return file_model - finally: - if temp_path: - temp_path.unlink(missing_ok=True) - - def _save(self, file: File, new: bool = False) -> None: - self._files_dir.mkdir(parents=True, exist_ok=True) - file_path = self._get_file_path(file.id, new=new) - content = file.model_dump_json() - file_path.write_text(content, encoding="utf-8") + """Retrieve file metadata and filesystem path.""" + file_metadata = self.retrieve(file_id) + file_path = self.get_file_path(file_id) + return file_metadata, file_path + return file_metadata, file_path + return file_metadata, file_path diff --git a/src/askui/chat/api/mcp_configs/dependencies.py b/src/askui/chat/api/mcp_configs/dependencies.py index 023b2bcb..e8423955 100644 --- a/src/askui/chat/api/mcp_configs/dependencies.py +++ b/src/askui/chat/api/mcp_configs/dependencies.py @@ -1,13 +1,14 @@ -from fastapi import Depends - -from askui.chat.api.dependencies import SettingsDep +from askui.chat.api.dependencies import SessionFactoryDep, SettingsDep from askui.chat.api.mcp_configs.service import McpConfigService from askui.chat.api.settings import Settings +from fastapi import Depends -def get_mcp_config_service(settings: Settings = SettingsDep) -> McpConfigService: +def get_mcp_config_service( + session_factory=SessionFactoryDep, settings: Settings = SettingsDep +) -> McpConfigService: """Get McpConfigService instance.""" - return McpConfigService(settings.data_dir, settings.mcp_configs) + return McpConfigService(session_factory, settings.data_dir, settings.mcp_configs) McpConfigServiceDep = Depends(get_mcp_config_service) diff --git a/src/askui/chat/api/mcp_configs/models.py b/src/askui/chat/api/mcp_configs/models.py index 09e35128..4d709266 100644 --- a/src/askui/chat/api/mcp_configs/models.py +++ b/src/askui/chat/api/mcp_configs/models.py @@ -1,56 +1,41 @@ -from typing import Literal - -from fastmcp.mcp_config import RemoteMCPServer, StdioMCPServer -from pydantic import BaseModel - -from askui.chat.api.models import McpConfigId, WorkspaceId, WorkspaceResource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id -from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven - -McpServer = StdioMCPServer | RemoteMCPServer - - -class McpConfigBase(BaseModel): - """Base MCP configuration model.""" - - name: str - mcp_server: McpServer - - -class McpConfigCreateParams(McpConfigBase): - """Parameters for creating an MCP configuration.""" - - -class McpConfigModifyParams(BaseModelWithNotGiven): - """Parameters for modifying an MCP configuration.""" - - name: str | NotGiven = NOT_GIVEN - mcp_server: McpServer | NotGiven = NOT_GIVEN - - -class McpConfig(McpConfigBase, WorkspaceResource): - """An MCP configuration that can be stored and managed.""" - - id: McpConfigId - object: Literal["mcp_config"] = "mcp_config" - created_at: UnixDatetime +"""MCP Config database model.""" + +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import McpConfigId +from askui.chat.api.mcp_configs.schemas import McpConfig +from sqlalchemy import JSON, Column, DateTime, String + + +class McpConfigModel(Base): + """MCP Config database model.""" + + __tablename__ = "mcp_configs" + id = Column(McpConfigId, primary_key=True) + workspace_id = Column(String(36), nullable=True, index=True) + created_at = Column(DateTime, nullable=False, index=True) + name = Column(String, nullable=False) + mcp_server = Column(JSON, nullable=False) + + def to_pydantic(self) -> McpConfig: + """Convert to Pydantic model.""" + data = { + "id": self.id, # Prefix is handled by the specialized type + "workspace_id": self.workspace_id, + "created_at": self.created_at, + "name": self.name, + "mcp_server": self.mcp_server, + } + return McpConfig.model_validate(data) @classmethod - def create( - cls, workspace_id: WorkspaceId, params: McpConfigCreateParams - ) -> "McpConfig": + def from_pydantic(cls, mcp_config: McpConfig) -> "McpConfigModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("mcpcnf"), - created_at=now(), - workspace_id=workspace_id, - **params.model_dump(), - ) - - def modify(self, params: McpConfigModifyParams) -> "McpConfig": - return McpConfig.model_validate( - { - **self.model_dump(), - **params.model_dump(), - } + id=mcp_config.id, + workspace_id=str(mcp_config.workspace_id) + if mcp_config.workspace_id + else None, + created_at=mcp_config.created_at, + name=mcp_config.name, + mcp_server=mcp_config.mcp_server.model_dump(), ) diff --git a/src/askui/chat/api/mcp_configs/router.py b/src/askui/chat/api/mcp_configs/router.py index 07f3b7ac..a9336524 100644 --- a/src/askui/chat/api/mcp_configs/router.py +++ b/src/askui/chat/api/mcp_configs/router.py @@ -1,10 +1,8 @@ from typing import Annotated -from fastapi import APIRouter, Header, status - from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.mcp_configs.dependencies import McpConfigServiceDep -from askui.chat.api.mcp_configs.models import ( +from askui.chat.api.mcp_configs.schemas import ( McpConfig, McpConfigCreateParams, McpConfigModifyParams, @@ -12,6 +10,7 @@ from askui.chat.api.mcp_configs.service import McpConfigService from askui.chat.api.models import McpConfigId, WorkspaceId from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, Header, status router = APIRouter(prefix="/mcp-configs", tags=["mcp-configs"]) diff --git a/src/askui/chat/api/mcp_configs/schemas.py b/src/askui/chat/api/mcp_configs/schemas.py new file mode 100644 index 00000000..09e35128 --- /dev/null +++ b/src/askui/chat/api/mcp_configs/schemas.py @@ -0,0 +1,56 @@ +from typing import Literal + +from fastmcp.mcp_config import RemoteMCPServer, StdioMCPServer +from pydantic import BaseModel + +from askui.chat.api.models import McpConfigId, WorkspaceId, WorkspaceResource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id +from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven + +McpServer = StdioMCPServer | RemoteMCPServer + + +class McpConfigBase(BaseModel): + """Base MCP configuration model.""" + + name: str + mcp_server: McpServer + + +class McpConfigCreateParams(McpConfigBase): + """Parameters for creating an MCP configuration.""" + + +class McpConfigModifyParams(BaseModelWithNotGiven): + """Parameters for modifying an MCP configuration.""" + + name: str | NotGiven = NOT_GIVEN + mcp_server: McpServer | NotGiven = NOT_GIVEN + + +class McpConfig(McpConfigBase, WorkspaceResource): + """An MCP configuration that can be stored and managed.""" + + id: McpConfigId + object: Literal["mcp_config"] = "mcp_config" + created_at: UnixDatetime + + @classmethod + def create( + cls, workspace_id: WorkspaceId, params: McpConfigCreateParams + ) -> "McpConfig": + return cls( + id=generate_time_ordered_id("mcpcnf"), + created_at=now(), + workspace_id=workspace_id, + **params.model_dump(), + ) + + def modify(self, params: McpConfigModifyParams) -> "McpConfig": + return McpConfig.model_validate( + { + **self.model_dump(), + **params.model_dump(), + } + ) diff --git a/src/askui/chat/api/mcp_configs/service.py b/src/askui/chat/api/mcp_configs/service.py index 4c376142..cdc098be 100644 --- a/src/askui/chat/api/mcp_configs/service.py +++ b/src/askui/chat/api/mcp_configs/service.py @@ -1,75 +1,90 @@ from pathlib import Path +from typing import Callable -from fastmcp.mcp_config import MCPConfig - -from askui.chat.api.mcp_configs.models import ( +from askui.chat.api.db.query_builder import QueryBuilder +from askui.chat.api.mcp_configs.models import McpConfigModel +from askui.chat.api.mcp_configs.schemas import ( McpConfig, McpConfigCreateParams, McpConfigId, McpConfigModifyParams, ) from askui.chat.api.models import WorkspaceId -from askui.chat.api.utils import build_workspace_filter_fn from askui.utils.api_utils import ( LIST_LIMIT_MAX, - ConflictError, ForbiddenError, LimitReachedError, ListQuery, ListResponse, NotFoundError, - list_resources, ) +from fastmcp.mcp_config import MCPConfig +from sqlalchemy.orm import Session class McpConfigService: - """Service for managing McpConfig resources with filesystem persistence.""" + """Service for managing McpConfig resources with SQLAlchemy persistence.""" - def __init__(self, base_dir: Path, seeds: list[McpConfig]) -> None: + def __init__( + self, + session_factory: Callable[[], Session], + base_dir: Path, + seeds: list[McpConfig], + ) -> None: + self._session_factory = session_factory self._base_dir = base_dir - self._mcp_configs_dir = base_dir / "mcp_configs" self._seeds = seeds - def _get_mcp_config_path( - self, mcp_config_id: McpConfigId, new: bool = False - ) -> Path: - mcp_config_path = self._mcp_configs_dir / f"{mcp_config_id}.json" - exists = mcp_config_path.exists() - if new and exists: - error_msg = f"MCP configuration {mcp_config_id} already exists" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"MCP configuration {mcp_config_id} not found" - raise NotFoundError(error_msg) - return mcp_config_path + def _to_pydantic(self, db_model: McpConfigModel) -> McpConfig: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_( self, workspace_id: WorkspaceId | None, query: ListQuery ) -> ListResponse[McpConfig]: - return list_resources( - self._mcp_configs_dir, - query, - McpConfig, - filter_fn=build_workspace_filter_fn(workspace_id, McpConfig), - ) + with self._session_factory() as session: + q = session.query(McpConfigModel) + + # Filter by workspace + if workspace_id is not None: + q = q.filter(McpConfigModel.workspace_id == str(workspace_id)) + else: + q = q.filter(McpConfigModel.workspace_id.is_(None)) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, McpConfigModel, query, McpConfigModel.created_at, McpConfigModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve( self, workspace_id: WorkspaceId | None, mcp_config_id: McpConfigId ) -> McpConfig: - try: - mcp_config_path = self._get_mcp_config_path(mcp_config_id) - mcp_config = McpConfig.model_validate_json(mcp_config_path.read_text()) + with self._session_factory() as session: + db_config = ( + session.query(McpConfigModel) + .filter(McpConfigModel.id == mcp_config_id) + .first() + ) + if not db_config: + error_msg = f"MCP configuration {mcp_config_id} not found" + raise NotFoundError(error_msg) + + # Check workspace access if not ( - mcp_config.workspace_id is None - or mcp_config.workspace_id == workspace_id + db_config.workspace_id is None + or db_config.workspace_id == str(workspace_id) ): error_msg = f"MCP configuration {mcp_config_id} not found" raise NotFoundError(error_msg) - except FileNotFoundError as e: - error_msg = f"MCP configuration {mcp_config_id} not found" - raise NotFoundError(error_msg) from e - else: - return mcp_config + + return self._to_pydantic(db_config) def retrieve_fast_mcp_config( self, workspace_id: WorkspaceId | None @@ -98,9 +113,12 @@ def create( self, workspace_id: WorkspaceId, params: McpConfigCreateParams ) -> McpConfig: self._check_limit(workspace_id) - mcp_config = McpConfig.create(workspace_id, params) - self._save(mcp_config, new=True) - return mcp_config + with self._session_factory() as session: + db_config = McpConfigModel.from_create_params(params, workspace_id) + session.add(db_config) + session.commit() + session.refresh(db_config) + return self._to_pydantic(db_config) def modify( self, @@ -108,13 +126,39 @@ def modify( mcp_config_id: McpConfigId, params: McpConfigModifyParams, ) -> McpConfig: - mcp_config = self.retrieve(workspace_id, mcp_config_id) - if mcp_config.workspace_id is None: - error_msg = f"Default MCP configuration {mcp_config_id} cannot be modified" - raise ForbiddenError(error_msg) - modified = mcp_config.modify(params) - self._save(modified) - return modified + with self._session_factory() as session: + db_config = ( + session.query(McpConfigModel) + .filter(McpConfigModel.id == mcp_config_id) + .first() + ) + if not db_config: + error_msg = f"MCP configuration {mcp_config_id} not found" + raise NotFoundError(error_msg) + + # Check workspace access + if not ( + db_config.workspace_id is None + or db_config.workspace_id == str(workspace_id) + ): + error_msg = f"MCP configuration {mcp_config_id} not found" + raise NotFoundError(error_msg) + + if db_config.workspace_id is None: + error_msg = ( + f"Default MCP configuration {mcp_config_id} cannot be modified" + ) + raise ForbiddenError(error_msg) + + # Update fields + if params.name is not None: + db_config.name = params.name + if params.mcp_server is not None: + db_config.mcp_server = params.mcp_server + + session.commit() + session.refresh(db_config) + return self._to_pydantic(db_config) def delete( self, @@ -122,37 +166,50 @@ def delete( mcp_config_id: McpConfigId, force: bool = False, ) -> None: - try: - mcp_config = self.retrieve(workspace_id, mcp_config_id) - if mcp_config.workspace_id is None and not force: + with self._session_factory() as session: + db_config = ( + session.query(McpConfigModel) + .filter(McpConfigModel.id == mcp_config_id) + .first() + ) + if not db_config: + error_msg = f"MCP configuration {mcp_config_id} not found" + if not force: + raise NotFoundError(error_msg) + return + + # Check workspace access + if not ( + db_config.workspace_id is None + or db_config.workspace_id == str(workspace_id) + ): + error_msg = f"MCP configuration {mcp_config_id} not found" + if not force: + raise NotFoundError(error_msg) + return + + if db_config.workspace_id is None and not force: error_msg = ( f"Default MCP configuration {mcp_config_id} cannot be deleted" ) raise ForbiddenError(error_msg) - self._get_mcp_config_path(mcp_config_id).unlink() - except FileNotFoundError as e: - error_msg = f"MCP configuration {mcp_config_id} not found" - if not force: - raise NotFoundError(error_msg) from e - except NotFoundError: - if not force: - raise - - def _save(self, mcp_config: McpConfig, new: bool = False) -> None: - self._mcp_configs_dir.mkdir(parents=True, exist_ok=True) - mcp_config_file = self._get_mcp_config_path(mcp_config.id, new=new) - mcp_config_file.write_text( - mcp_config.model_dump_json( - exclude_unset=True, exclude_none=True, exclude_defaults=True - ), - encoding="utf-8", - ) + + session.delete(db_config) + session.commit() def seed(self) -> None: """Seed the MCP configuration service with default MCP configurations.""" - for seed in self._seeds: - try: - self.delete(None, seed.id, force=True) - self._save(seed, new=True) - except ConflictError: # noqa: PERF203 - self._save(seed) + with self._session_factory() as session: + for seed in self._seeds: + # Check if config already exists + existing_config = ( + session.query(McpConfigModel) + .filter(McpConfigModel.id == seed.id) + .first() + ) + if not existing_config: + db_config = McpConfigModel.from_pydantic(seed) + session.add(db_config) + + session.commit() + session.commit() diff --git a/src/askui/chat/api/messages/chat_history_manager.py b/src/askui/chat/api/messages/chat_history_manager.py index 9642257c..2888616a 100644 --- a/src/askui/chat/api/messages/chat_history_manager.py +++ b/src/askui/chat/api/messages/chat_history_manager.py @@ -1,6 +1,5 @@ from anthropic.types.beta import BetaTextBlockParam, BetaToolUnionParam - -from askui.chat.api.messages.models import Message, MessageCreateParams +from askui.chat.api.messages.schemas import Message, MessageCreateParams from askui.chat.api.messages.service import MessageService from askui.chat.api.messages.translator import MessageTranslator from askui.chat.api.models import ThreadId diff --git a/src/askui/chat/api/messages/dependencies.py b/src/askui/chat/api/messages/dependencies.py index e22ea940..7e7c553a 100644 --- a/src/askui/chat/api/messages/dependencies.py +++ b/src/askui/chat/api/messages/dependencies.py @@ -1,8 +1,4 @@ -from pathlib import Path - -from fastapi import Depends - -from askui.chat.api.dependencies import WorkspaceDirDep +from askui.chat.api.dependencies import SessionFactoryDep from askui.chat.api.files.dependencies import FileServiceDep from askui.chat.api.files.service import FileService from askui.chat.api.messages.chat_history_manager import ChatHistoryManager @@ -12,13 +8,12 @@ SimpleTruncationStrategyFactory, TruncationStrategyFactory, ) +from fastapi import Depends -def get_message_service( - workspace_dir: Path = WorkspaceDirDep, -) -> MessageService: - """Get MessagePersistedService instance.""" - return MessageService(workspace_dir) +def get_message_service(session_factory=SessionFactoryDep) -> MessageService: + """Get MessageService instance.""" + return MessageService(session_factory) MessageServiceDep = Depends(get_message_service) @@ -53,3 +48,4 @@ def get_chat_history_manager( ChatHistoryManagerDep = Depends(get_chat_history_manager) +ChatHistoryManagerDep = Depends(get_chat_history_manager) diff --git a/src/askui/chat/api/messages/models.py b/src/askui/chat/api/messages/models.py index 346fff7e..046ed6e1 100644 --- a/src/askui/chat/api/messages/models.py +++ b/src/askui/chat/api/messages/models.py @@ -1,95 +1,62 @@ -from typing import Literal - -from pydantic import BaseModel - -from askui.chat.api.models import AssistantId, FileId, MessageId, RunId, ThreadId -from askui.models.shared.agent_message_param import ( - Base64ImageSourceParam, - BetaRedactedThinkingBlock, - BetaThinkingBlock, - CacheControlEphemeralParam, - StopReason, - TextBlockParam, - ToolUseBlockParam, - UrlImageSourceParam, -) -from askui.utils.api_utils import Resource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id - - -class BetaFileDocumentSourceParam(BaseModel): - file_id: str - type: Literal["file"] = "file" - - -Source = BetaFileDocumentSourceParam - - -class RequestDocumentBlockParam(BaseModel): - source: Source - type: Literal["document"] = "document" - cache_control: CacheControlEphemeralParam | None = None - - -class FileImageSourceParam(BaseModel): - """Image source that references a saved file.""" - - id: FileId - type: Literal["file"] = "file" - - -class ImageBlockParam(BaseModel): - source: Base64ImageSourceParam | UrlImageSourceParam | FileImageSourceParam - type: Literal["image"] = "image" - cache_control: CacheControlEphemeralParam | None = None - - -class ToolResultBlockParam(BaseModel): - tool_use_id: str - type: Literal["tool_result"] = "tool_result" - cache_control: CacheControlEphemeralParam | None = None - content: str | list[TextBlockParam | ImageBlockParam] - is_error: bool = False - - -ContentBlockParam = ( - ImageBlockParam - | TextBlockParam - | ToolResultBlockParam - | ToolUseBlockParam - | BetaThinkingBlock - | BetaRedactedThinkingBlock - | RequestDocumentBlockParam -) - - -class MessageParam(BaseModel): - role: Literal["user", "assistant"] - content: str | list[ContentBlockParam] - stop_reason: StopReason | None = None - - -class MessageBase(MessageParam): - assistant_id: AssistantId | None = None - run_id: RunId | None = None - - -class MessageCreateParams(MessageBase): - pass - - -class Message(MessageBase, Resource): - id: MessageId - object: Literal["thread.message"] = "thread.message" - created_at: UnixDatetime - thread_id: ThreadId +"""Message database model.""" + +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import AssistantId, MessageId, RunId, ThreadId +from askui.chat.api.messages.schemas import Message +from sqlalchemy import JSON, Column, DateTime, ForeignKey, String + + +class MessageModel(Base): + """Message database model.""" + + __tablename__ = "messages" + id = Column(MessageId, primary_key=True) + thread_id = Column( + ThreadId, + ForeignKey("threads.id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + created_at = Column(DateTime, nullable=False, index=True) + assistant_id = Column( + AssistantId, + ForeignKey("assistants.id", ondelete="SET NULL"), + nullable=True, + index=True, + ) + run_id = Column( + RunId, + ForeignKey("runs.id", ondelete="SET NULL"), + nullable=True, + index=True, + ) + role = Column(String, nullable=False) + content = Column(JSON, nullable=False) + stop_reason = Column(String, nullable=True) + + def to_pydantic(self) -> Message: + """Convert to Pydantic model.""" + return Message( + id=self.id, # Prefix is handled by the specialized type + thread_id=self.thread_id, + created_at=self.created_at, + assistant_id=self.assistant_id, + run_id=self.run_id, + role=self.role, + content=self.content, + stop_reason=self.stop_reason, + ) @classmethod - def create(cls, thread_id: ThreadId, params: MessageCreateParams) -> "Message": + def from_pydantic(cls, message: Message) -> "MessageModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("msg"), - created_at=now(), - thread_id=thread_id, - **params.model_dump(), + id=message.id, + thread_id=message.thread_id, + created_at=message.created_at, + assistant_id=message.assistant_id, + run_id=message.run_id, + role=message.role, + content=message.content, + stop_reason=message.stop_reason, ) diff --git a/src/askui/chat/api/messages/router.py b/src/askui/chat/api/messages/router.py index 4276950a..9a4a5844 100644 --- a/src/askui/chat/api/messages/router.py +++ b/src/askui/chat/api/messages/router.py @@ -1,13 +1,12 @@ -from fastapi import APIRouter, status - from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.messages.dependencies import MessageServiceDep -from askui.chat.api.messages.models import Message, MessageCreateParams +from askui.chat.api.messages.schemas import Message, MessageCreateParams from askui.chat.api.messages.service import MessageService from askui.chat.api.models import MessageId, ThreadId from askui.chat.api.threads.dependencies import ThreadFacadeDep from askui.chat.api.threads.facade import ThreadFacade from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, status router = APIRouter(prefix="/threads/{thread_id}/messages", tags=["messages"]) diff --git a/src/askui/chat/api/messages/schemas.py b/src/askui/chat/api/messages/schemas.py new file mode 100644 index 00000000..346fff7e --- /dev/null +++ b/src/askui/chat/api/messages/schemas.py @@ -0,0 +1,95 @@ +from typing import Literal + +from pydantic import BaseModel + +from askui.chat.api.models import AssistantId, FileId, MessageId, RunId, ThreadId +from askui.models.shared.agent_message_param import ( + Base64ImageSourceParam, + BetaRedactedThinkingBlock, + BetaThinkingBlock, + CacheControlEphemeralParam, + StopReason, + TextBlockParam, + ToolUseBlockParam, + UrlImageSourceParam, +) +from askui.utils.api_utils import Resource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id + + +class BetaFileDocumentSourceParam(BaseModel): + file_id: str + type: Literal["file"] = "file" + + +Source = BetaFileDocumentSourceParam + + +class RequestDocumentBlockParam(BaseModel): + source: Source + type: Literal["document"] = "document" + cache_control: CacheControlEphemeralParam | None = None + + +class FileImageSourceParam(BaseModel): + """Image source that references a saved file.""" + + id: FileId + type: Literal["file"] = "file" + + +class ImageBlockParam(BaseModel): + source: Base64ImageSourceParam | UrlImageSourceParam | FileImageSourceParam + type: Literal["image"] = "image" + cache_control: CacheControlEphemeralParam | None = None + + +class ToolResultBlockParam(BaseModel): + tool_use_id: str + type: Literal["tool_result"] = "tool_result" + cache_control: CacheControlEphemeralParam | None = None + content: str | list[TextBlockParam | ImageBlockParam] + is_error: bool = False + + +ContentBlockParam = ( + ImageBlockParam + | TextBlockParam + | ToolResultBlockParam + | ToolUseBlockParam + | BetaThinkingBlock + | BetaRedactedThinkingBlock + | RequestDocumentBlockParam +) + + +class MessageParam(BaseModel): + role: Literal["user", "assistant"] + content: str | list[ContentBlockParam] + stop_reason: StopReason | None = None + + +class MessageBase(MessageParam): + assistant_id: AssistantId | None = None + run_id: RunId | None = None + + +class MessageCreateParams(MessageBase): + pass + + +class Message(MessageBase, Resource): + id: MessageId + object: Literal["thread.message"] = "thread.message" + created_at: UnixDatetime + thread_id: ThreadId + + @classmethod + def create(cls, thread_id: ThreadId, params: MessageCreateParams) -> "Message": + return cls( + id=generate_time_ordered_id("msg"), + created_at=now(), + thread_id=thread_id, + **params.model_dump(), + ) diff --git a/src/askui/chat/api/messages/service.py b/src/askui/chat/api/messages/service.py index 1d2a4781..056dfec9 100644 --- a/src/askui/chat/api/messages/service.py +++ b/src/askui/chat/api/messages/service.py @@ -1,83 +1,99 @@ +"""Message service with SQLAlchemy persistence.""" + from pathlib import Path -from typing import Iterator +from typing import Callable, Iterator -from askui.chat.api.messages.models import Message, MessageCreateParams +from askui.chat.api.db.query_builder import QueryBuilder +from askui.chat.api.messages.models import MessageModel +from askui.chat.api.messages.schemas import Message, MessageCreateParams from askui.chat.api.models import MessageId, ThreadId -from askui.utils.api_utils import ( - LIST_LIMIT_DEFAULT, - ConflictError, - ListOrder, - ListQuery, - ListResponse, - NotFoundError, - list_resources, -) +from askui.utils.api_utils import ListQuery, ListResponse, NotFoundError +from sqlalchemy.orm import Session class MessageService: - def __init__(self, base_dir: Path) -> None: - self._base_dir = base_dir + """Service for managing Message resources with SQLAlchemy persistence.""" - def get_messages_dir(self, thread_id: ThreadId) -> Path: - return self._base_dir / "messages" / thread_id - - def _get_message_path( - self, thread_id: ThreadId, message_id: MessageId, new: bool = False - ) -> Path: - message_path = self.get_messages_dir(thread_id) / f"{message_id}.json" - exists = message_path.exists() - if new and exists: - error_msg = f"Message {message_id} already exists in thread {thread_id}" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"Message {message_id} not found in thread {thread_id}" - raise NotFoundError(error_msg) - return message_path + def __init__(self, session_factory: Callable[[], Session]) -> None: + self._session_factory = session_factory - def create(self, thread_id: ThreadId, params: MessageCreateParams) -> Message: - new_message = Message.create(thread_id, params) - self._save(new_message, new=True) - return new_message + def _to_pydantic(self, db_model: MessageModel) -> Message: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_(self, thread_id: ThreadId, query: ListQuery) -> ListResponse[Message]: - messages_dir = self.get_messages_dir(thread_id) - return list_resources(messages_dir, query, Message) - - def iter( - self, - thread_id: ThreadId, - order: ListOrder = "asc", - batch_size: int = LIST_LIMIT_DEFAULT, - ) -> Iterator[Message]: - has_more = True - last_id: str | None = None - while has_more: - list_messages_response = self.list_( - thread_id=thread_id, - query=ListQuery(limit=batch_size, order=order, after=last_id), + """List messages for a thread with pagination.""" + with self._session_factory() as session: + q = session.query(MessageModel).filter(MessageModel.thread_id == thread_id) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, MessageModel, query, MessageModel.created_at, MessageModel.id ) - has_more = list_messages_response.has_more - last_id = list_messages_response.last_id - for msg in list_messages_response.data: - yield msg + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve(self, thread_id: ThreadId, message_id: MessageId) -> Message: - try: - message_file = self._get_message_path(thread_id, message_id) - return Message.model_validate_json(message_file.read_text(encoding="utf-8")) - except FileNotFoundError as e: - error_msg = f"Message {message_id} not found in thread {thread_id}" - raise NotFoundError(error_msg) from e + """Retrieve a message by ID.""" + with self._session_factory() as session: + db_message = ( + session.query(MessageModel) + .filter( + MessageModel.thread_id == thread_id, + MessageModel.id == message_id, + ) + .first() + ) + if not db_message: + error_msg = f"Message {message_id} not found" + raise NotFoundError(error_msg) + return self._to_pydantic(db_message) + + def create(self, thread_id: ThreadId, params: MessageCreateParams) -> Message: + """Create a new message.""" + with self._session_factory() as session: + db_message = MessageModel.from_create_params(params, thread_id) + session.add(db_message) + session.commit() + session.refresh(db_message) + + return self._to_pydantic(db_message) def delete(self, thread_id: ThreadId, message_id: MessageId) -> None: - try: - self._get_message_path(thread_id, message_id).unlink() - except FileNotFoundError as e: - error_msg = f"Message {message_id} not found in thread {thread_id}" - raise NotFoundError(error_msg) from e - - def _save(self, message: Message, new: bool = False) -> None: - messages_dir = self.get_messages_dir(message.thread_id) - messages_dir.mkdir(parents=True, exist_ok=True) - message_file = self._get_message_path(message.thread_id, message.id, new=new) - message_file.write_text(message.model_dump_json(), encoding="utf-8") + """Delete a message.""" + with self._session_factory() as session: + db_message = ( + session.query(MessageModel) + .filter( + MessageModel.thread_id == thread_id, + MessageModel.id == message_id, + ) + .first() + ) + if not db_message: + error_msg = f"Message {message_id} not found" + raise NotFoundError(error_msg) + + session.delete(db_message) + session.commit() + + def get_messages_dir(self, thread_id: ThreadId) -> Path: + """Get messages directory for a thread (for backward compatibility).""" + return Path.cwd() / "chat" / "messages" / thread_id + + def list_messages(self, thread_id: ThreadId) -> Iterator[Message]: + """List all messages for a thread (for backward compatibility).""" + with self._session_factory() as session: + db_messages = ( + session.query(MessageModel) + .filter(MessageModel.thread_id == thread_id) + .order_by(MessageModel.created_at) + .all() + ) + for db_message in db_messages: + yield self._to_pydantic(db_message) diff --git a/src/askui/chat/api/messages/translator.py b/src/askui/chat/api/messages/translator.py index 04f448ba..5701693b 100644 --- a/src/askui/chat/api/messages/translator.py +++ b/src/askui/chat/api/messages/translator.py @@ -1,7 +1,5 @@ -from PIL import Image - from askui.chat.api.files.service import FileService -from askui.chat.api.messages.models import ( +from askui.chat.api.messages.schemas import ( ContentBlockParam, FileImageSourceParam, ImageBlockParam, @@ -11,11 +9,7 @@ ) from askui.data_extractor import DataExtractor from askui.models.models import ModelName -from askui.models.shared.agent_message_param import ( - Base64ImageSourceParam, - TextBlockParam, - UrlImageSourceParam, -) +from askui.models.shared.agent_message_param import Base64ImageSourceParam from askui.models.shared.agent_message_param import ( ContentBlockParam as AnthropicContentBlockParam, ) @@ -25,12 +19,15 @@ from askui.models.shared.agent_message_param import ( MessageParam as AnthropicMessageParam, ) +from askui.models.shared.agent_message_param import TextBlockParam from askui.models.shared.agent_message_param import ( ToolResultBlockParam as AnthropicToolResultBlockParam, ) +from askui.models.shared.agent_message_param import UrlImageSourceParam from askui.utils.excel_utils import OfficeDocumentSource from askui.utils.image_utils import ImageSource, image_to_base64 from askui.utils.source_utils import Source, load_source +from PIL import Image class RequestDocumentBlockParamTranslator: diff --git a/src/askui/chat/api/migrations/models.py b/src/askui/chat/api/migrations/models.py new file mode 100644 index 00000000..05dfb71a --- /dev/null +++ b/src/askui/chat/api/migrations/models.py @@ -0,0 +1,12 @@ +"""Migration version database model.""" + +from askui.chat.api.db.base import Base +from sqlalchemy import Column, DateTime, Integer + + +class MigrationVersionModel(Base): + """Migration version tracking model.""" + + __tablename__ = "migration_version" + version = Column(Integer, primary_key=True) + applied_at = Column(DateTime, nullable=False) diff --git a/src/askui/chat/api/runs/dependencies.py b/src/askui/chat/api/runs/dependencies.py index 759fbf73..7ed80a04 100644 --- a/src/askui/chat/api/runs/dependencies.py +++ b/src/askui/chat/api/runs/dependencies.py @@ -1,16 +1,13 @@ -from pathlib import Path - -from fastapi import Depends - from askui.chat.api.assistants.dependencies import AssistantServiceDep from askui.chat.api.assistants.service import AssistantService -from askui.chat.api.dependencies import SettingsDep, WorkspaceDirDep +from askui.chat.api.dependencies import SessionFactoryDep, SettingsDep from askui.chat.api.mcp_clients.dependencies import McpClientManagerManagerDep from askui.chat.api.mcp_clients.manager import McpClientManagerManager from askui.chat.api.messages.chat_history_manager import ChatHistoryManager from askui.chat.api.messages.dependencies import ChatHistoryManagerDep -from askui.chat.api.runs.models import RunListQuery +from askui.chat.api.runs.schemas import RunListQuery from askui.chat.api.settings import Settings +from fastapi import Depends from .service import RunService @@ -18,14 +15,14 @@ def get_runs_service( - workspace_dir: Path = WorkspaceDirDep, + session_factory=SessionFactoryDep, assistant_service: AssistantService = AssistantServiceDep, chat_history_manager: ChatHistoryManager = ChatHistoryManagerDep, mcp_client_manager_manager: McpClientManagerManager = McpClientManagerManagerDep, settings: Settings = SettingsDep, ) -> RunService: return RunService( - base_dir=workspace_dir, + session_factory=session_factory, assistant_service=assistant_service, mcp_client_manager_manager=mcp_client_manager_manager, chat_history_manager=chat_history_manager, diff --git a/src/askui/chat/api/runs/events/models.py b/src/askui/chat/api/runs/events/models.py new file mode 100644 index 00000000..66b08dc4 --- /dev/null +++ b/src/askui/chat/api/runs/events/models.py @@ -0,0 +1,45 @@ +"""Event database model.""" + +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import RunId, ThreadId +from askui.chat.api.runs.events.events import Event +from sqlalchemy import JSON, Column, DateTime, ForeignKey, Index, Integer, String + + +class EventModel(Base): + """Event database model.""" + + __tablename__ = "events" + id = Column(Integer, primary_key=True, autoincrement=True) + run_id = Column( + RunId, + ForeignKey("runs.id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + thread_id = Column(ThreadId, nullable=False, index=True) + sequence_num = Column(Integer, nullable=False, index=True) + event_type = Column(String, nullable=False) + event_data = Column(JSON, nullable=False) + created_at = Column(DateTime, nullable=False, index=True) + + __table_args__ = (Index("idx_events_run_sequence", "run_id", "sequence_num"),) + + def to_pydantic(self) -> Event: + """Convert to Pydantic model.""" + # Parse the event data JSON to create Event object + return Event.model_validate_json(self.event_data) + + @classmethod + def from_pydantic( + cls, event: Event, run_id: str, thread_id: str, sequence_num: int + ) -> "EventModel": + """Create from Pydantic model.""" + return cls( + run_id=run_id, + thread_id=thread_id, + sequence_num=sequence_num, + event_type=event.event, + event_data=event.model_dump_json(), + created_at=event.created_at, + ) diff --git a/src/askui/chat/api/runs/models.py b/src/askui/chat/api/runs/models.py index 96220ffd..3229881f 100644 --- a/src/askui/chat/api/runs/models.py +++ b/src/askui/chat/api/runs/models.py @@ -1,111 +1,67 @@ -from dataclasses import dataclass -from datetime import timedelta -from typing import Annotated, Literal - -from fastapi import Query -from pydantic import BaseModel, computed_field - -from askui.chat.api.models import AssistantId, RunId, ThreadId -from askui.chat.api.threads.models import ThreadCreateParams -from askui.utils.api_utils import ListQuery, Resource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id - -RunStatus = Literal[ - "queued", - "in_progress", - "completed", - "cancelling", - "cancelled", - "failed", - "expired", -] - - -class RunError(BaseModel): - """Error information for a failed run.""" - - message: str - code: Literal["server_error", "rate_limit_exceeded", "invalid_prompt"] - - -class RunBase(BaseModel): - """Base run model.""" - - assistant_id: AssistantId - - -class RunCreateParams(RunBase): - """Parameters for creating a run.""" - - stream: bool = False - - -class ThreadAndRunCreateParams(RunCreateParams): - thread: ThreadCreateParams - - -class Run(RunBase, Resource): - """A run execution within a thread.""" - - id: RunId - object: Literal["thread.run"] = "thread.run" - thread_id: ThreadId - created_at: UnixDatetime - expires_at: UnixDatetime - started_at: UnixDatetime | None = None - completed_at: UnixDatetime | None = None - failed_at: UnixDatetime | None = None - cancelled_at: UnixDatetime | None = None - tried_cancelling_at: UnixDatetime | None = None - last_error: RunError | None = None +"""Run database model.""" + +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import AssistantId, RunId, ThreadId +from askui.chat.api.runs.schemas import Run +from sqlalchemy import JSON, Column, DateTime, ForeignKey + + +class RunModel(Base): + """Run database model.""" + + __tablename__ = "runs" + id = Column(RunId, primary_key=True) + thread_id = Column( + ThreadId, + ForeignKey("threads.id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + assistant_id = Column( + AssistantId, + ForeignKey("assistants.id", ondelete="SET NULL"), + nullable=True, + index=True, + ) + created_at = Column(DateTime, nullable=False, index=True) + started_at = Column(DateTime, nullable=True) + completed_at = Column(DateTime, nullable=True) + failed_at = Column(DateTime, nullable=True) + cancelled_at = Column(DateTime, nullable=True) + tried_cancelling_at = Column(DateTime, nullable=True) + expires_at = Column(DateTime, nullable=True) + last_error = Column(JSON, nullable=True) + + def to_pydantic(self) -> Run: + """Convert to Pydantic model.""" + data = { + "id": self.id, # Prefix is handled by the specialized type + "thread_id": self.thread_id, + "assistant_id": self.assistant_id, + "created_at": self.created_at, + "started_at": self.started_at, + "completed_at": self.completed_at, + "failed_at": self.failed_at, + "cancelled_at": self.cancelled_at, + "tried_cancelling_at": self.tried_cancelling_at, + "expires_at": self.expires_at, + "last_error": self.last_error, + } + return Run.model_validate(data) @classmethod - def create(cls, thread_id: ThreadId, params: RunCreateParams) -> "Run": + def from_pydantic(cls, run: Run) -> "RunModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("run"), - thread_id=thread_id, - created_at=now(), - expires_at=now() + timedelta(minutes=10), - **params.model_dump(exclude={"stream"}), + id=run.id, + thread_id=run.thread_id, + assistant_id=run.assistant_id, + created_at=run.created_at, + started_at=run.started_at, + completed_at=run.completed_at, + failed_at=run.failed_at, + cancelled_at=run.cancelled_at, + tried_cancelling_at=run.tried_cancelling_at, + expires_at=run.expires_at, + last_error=run.last_error.model_dump() if run.last_error else None, ) - - @computed_field # type: ignore[prop-decorator] - @property - def status(self) -> RunStatus: - if self.cancelled_at: - return "cancelled" - if self.failed_at: - return "failed" - if self.completed_at: - return "completed" - if self.expires_at and self.expires_at < now(): - return "expired" - if self.tried_cancelling_at: - return "cancelling" - if self.started_at: - return "in_progress" - return "queued" - - def start(self) -> None: - self.started_at = now() - self.expires_at = now() + timedelta(minutes=10) - - def ping(self) -> None: - self.expires_at = now() + timedelta(minutes=10) - - def complete(self) -> None: - self.completed_at = now() - - def cancel(self) -> None: - self.cancelled_at = now() - - def fail(self, error: RunError) -> None: - self.failed_at = now() - self.last_error = error - - -@dataclass(kw_only=True) -class RunListQuery(ListQuery): - thread: Annotated[ThreadId | None, Query()] = None - status: Annotated[list[RunStatus] | None, Query()] = None diff --git a/src/askui/chat/api/runs/router.py b/src/askui/chat/api/runs/router.py index bca81eb2..992aefef 100644 --- a/src/askui/chat/api/runs/router.py +++ b/src/askui/chat/api/runs/router.py @@ -1,28 +1,17 @@ from collections.abc import AsyncGenerator from typing import Annotated -from fastapi import ( - APIRouter, - BackgroundTasks, - Depends, - Header, - Path, - Query, - Response, - status, -) -from fastapi.responses import JSONResponse, StreamingResponse -from pydantic import BaseModel - -from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.models import RunId, ThreadId, WorkspaceId -from askui.chat.api.runs.models import RunCreateParams +from askui.chat.api.runs.schemas import RunCreateParams from askui.chat.api.threads.dependencies import ThreadFacadeDep from askui.chat.api.threads.facade import ThreadFacade -from askui.utils.api_utils import ListQuery, ListResponse +from askui.utils.api_utils import ListResponse +from fastapi import APIRouter, BackgroundTasks, Header, Path, Query, Response, status +from fastapi.responses import JSONResponse, StreamingResponse +from pydantic import BaseModel from .dependencies import RunListQueryDep, RunServiceDep -from .models import Run, RunListQuery, ThreadAndRunCreateParams +from .schemas import Run, RunListQuery, ThreadAndRunCreateParams from .service import RunService router = APIRouter(tags=["runs"]) diff --git a/src/askui/chat/api/runs/runner/runner.py b/src/askui/chat/api/runs/runner/runner.py index ae0c703b..e52a5170 100644 --- a/src/askui/chat/api/runs/runner/runner.py +++ b/src/askui/chat/api/runs/runner/runner.py @@ -5,8 +5,6 @@ from anthropic.types.beta import BetaCacheControlEphemeralParam, BetaTextBlockParam from anyio.abc import ObjectStream -from asyncer import asyncify, syncify - from askui.chat.api.assistants.models import Assistant from askui.chat.api.mcp_clients.manager import McpClientManagerManager from askui.chat.api.messages.chat_history_manager import ChatHistoryManager @@ -21,7 +19,7 @@ from askui.chat.api.runs.events.message_events import MessageEvent from askui.chat.api.runs.events.run_events import RunEvent from askui.chat.api.runs.events.service import RetrieveRunService -from askui.chat.api.runs.models import Run, RunError +from askui.chat.api.runs.schemas import Run, RunError from askui.chat.api.settings import Settings from askui.custom_agent import CustomAgent from askui.models.models import ModelName @@ -30,6 +28,7 @@ from askui.models.shared.settings import ActSettings, MessageSettings from askui.models.shared.tools import ToolCollection from askui.prompts.system import caesr_system_prompt +from asyncer import asyncify, syncify logger = logging.getLogger(__name__) diff --git a/src/askui/chat/api/runs/schemas.py b/src/askui/chat/api/runs/schemas.py new file mode 100644 index 00000000..a39b4667 --- /dev/null +++ b/src/askui/chat/api/runs/schemas.py @@ -0,0 +1,110 @@ +from dataclasses import dataclass +from datetime import timedelta +from typing import Annotated, Literal + +from askui.chat.api.models import AssistantId, RunId, ThreadId +from askui.chat.api.threads.schemas import ThreadCreateParams +from askui.utils.api_utils import ListQuery, Resource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id +from fastapi import Query +from pydantic import BaseModel, computed_field + +RunStatus = Literal[ + "queued", + "in_progress", + "completed", + "cancelling", + "cancelled", + "failed", + "expired", +] + + +class RunError(BaseModel): + """Error information for a failed run.""" + + message: str + code: Literal["server_error", "rate_limit_exceeded", "invalid_prompt"] + + +class RunBase(BaseModel): + """Base run model.""" + + assistant_id: AssistantId + + +class RunCreateParams(RunBase): + """Parameters for creating a run.""" + + stream: bool = False + + +class ThreadAndRunCreateParams(RunCreateParams): + thread: ThreadCreateParams + + +class Run(RunBase, Resource): + """A run execution within a thread.""" + + id: RunId + object: Literal["thread.run"] = "thread.run" + thread_id: ThreadId + created_at: UnixDatetime + expires_at: UnixDatetime + started_at: UnixDatetime | None = None + completed_at: UnixDatetime | None = None + failed_at: UnixDatetime | None = None + cancelled_at: UnixDatetime | None = None + tried_cancelling_at: UnixDatetime | None = None + last_error: RunError | None = None + + @classmethod + def create(cls, thread_id: ThreadId, params: RunCreateParams) -> "Run": + return cls( + id=generate_time_ordered_id("run"), + thread_id=thread_id, + created_at=now(), + expires_at=now() + timedelta(minutes=10), + **params.model_dump(exclude={"stream"}), + ) + + @computed_field # type: ignore[prop-decorator] + @property + def status(self) -> RunStatus: + if self.cancelled_at: + return "cancelled" + if self.failed_at: + return "failed" + if self.completed_at: + return "completed" + if self.expires_at and self.expires_at < now(): + return "expired" + if self.tried_cancelling_at: + return "cancelling" + if self.started_at: + return "in_progress" + return "queued" + + def start(self) -> None: + self.started_at = now() + self.expires_at = now() + timedelta(minutes=10) + + def ping(self) -> None: + self.expires_at = now() + timedelta(minutes=10) + + def complete(self) -> None: + self.completed_at = now() + + def cancel(self) -> None: + self.cancelled_at = now() + + def fail(self, error: RunError) -> None: + self.failed_at = now() + self.last_error = error + + +@dataclass(kw_only=True) +class RunListQuery(ListQuery): + thread: Annotated[ThreadId | None, Query()] = None + status: Annotated[list[RunStatus] | None, Query()] = None diff --git a/src/askui/chat/api/runs/service.py b/src/askui/chat/api/runs/service.py index 76a442b1..5faf62b5 100644 --- a/src/askui/chat/api/runs/service.py +++ b/src/askui/chat/api/runs/service.py @@ -4,72 +4,51 @@ from typing import Callable import anyio -from typing_extensions import override - from askui.chat.api.assistants.service import AssistantService +from askui.chat.api.db.query_builder import QueryBuilder from askui.chat.api.mcp_clients.manager import McpClientManagerManager from askui.chat.api.messages.chat_history_manager import ChatHistoryManager from askui.chat.api.models import RunId, ThreadId, WorkspaceId from askui.chat.api.runs.events.events import DoneEvent, ErrorEvent, Event, RunEvent from askui.chat.api.runs.events.service import EventService -from askui.chat.api.runs.models import Run, RunCreateParams, RunListQuery +from askui.chat.api.runs.models import RunModel from askui.chat.api.runs.runner.runner import Runner, RunnerRunService +from askui.chat.api.runs.schemas import Run, RunCreateParams, RunListQuery from askui.chat.api.settings import Settings -from askui.utils.api_utils import ( - ConflictError, - ListResponse, - NotFoundError, - list_resources, -) - - -def _build_run_filter_fn(query: RunListQuery) -> Callable[[Run], bool]: - def filter_fn(run: Run) -> bool: - return (query.thread is None or run.thread_id == query.thread) and ( - query.status is None or run.status in query.status - ) - - return filter_fn +from askui.utils.api_utils import ConflictError, ListQuery, ListResponse, NotFoundError +from sqlalchemy.orm import Session +from typing_extensions import override class RunService(RunnerRunService): - """Service for managing Run resources with filesystem persistence.""" + """Service for managing Run resources with SQLAlchemy persistence.""" def __init__( self, - base_dir: Path, + session_factory: Callable[[], Session], assistant_service: AssistantService, mcp_client_manager_manager: McpClientManagerManager, chat_history_manager: ChatHistoryManager, settings: Settings, ) -> None: - self._base_dir = base_dir + self._session_factory = session_factory self._assistant_service = assistant_service self._mcp_client_manager_manager = mcp_client_manager_manager self._chat_history_manager = chat_history_manager self._settings = settings - self._event_service = EventService(base_dir, self) - - def get_runs_dir(self, thread_id: ThreadId) -> Path: - return self._base_dir / "runs" / thread_id - - def _get_run_path( - self, thread_id: ThreadId, run_id: RunId, new: bool = False - ) -> Path: - run_path = self.get_runs_dir(thread_id) / f"{run_id}.json" - exists = run_path.exists() - if new and exists: - error_msg = f"Run {run_id} already exists in thread {thread_id}" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"Run {run_id} not found in thread {thread_id}" - raise NotFoundError(error_msg) - return run_path + self._event_service = EventService(Path.cwd() / "chat", self) + + def _to_pydantic(self, db_model: RunModel) -> Run: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def _create(self, thread_id: ThreadId, params: RunCreateParams) -> Run: - run = Run.create(thread_id, params) - self.save(run, new=True) - return run + with self._session_factory() as session: + db_run = RunModel.from_create_params(params, thread_id) + session.add(db_run) + session.commit() + session.refresh(db_run) + return self._to_pydantic(db_run) async def create( self, @@ -137,12 +116,19 @@ async def run_runner() -> None: @override def retrieve(self, thread_id: ThreadId, run_id: RunId) -> Run: - try: - run_file = self._get_run_path(thread_id, run_id) - return Run.model_validate_json(run_file.read_text()) - except FileNotFoundError as e: - error_msg = f"Run {run_id} not found in thread {thread_id}" - raise NotFoundError(error_msg) from e + with self._session_factory() as session: + db_run = ( + session.query(RunModel) + .filter( + RunModel.id == run_id, + RunModel.thread_id == thread_id, + ) + .first() + ) + if not db_run: + error_msg = f"Run {run_id} not found in thread {thread_id}" + raise NotFoundError(error_msg) + return self._to_pydantic(db_run) async def retrieve_stream( self, thread_id: ThreadId, run_id: RunId @@ -152,31 +138,110 @@ async def retrieve_stream( yield event def list_(self, query: RunListQuery) -> ListResponse[Run]: - if query.thread: - runs_dir = self.get_runs_dir(query.thread) - pattern = "*.json" - else: - runs_dir = self._base_dir / "runs" - pattern = "*/*.json" - return list_resources( - runs_dir, - query, - Run, - filter_fn=_build_run_filter_fn(query), - pattern=pattern, - ) + with self._session_factory() as session: + q = session.query(RunModel) + + # Filter by thread if specified + if query.thread: + q = q.filter(RunModel.thread_id == query.thread) + + # Filter by status if specified + if query.status: + q = q.filter(RunModel.status.in_(query.status)) + + # Convert to ListQuery for QueryBuilder + list_query = ListQuery( + limit=query.limit, + order=query.order, + after=query.after, + before=query.before, + ) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, RunModel, list_query, RunModel.created_at, RunModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def cancel(self, thread_id: ThreadId, run_id: RunId) -> Run: - run = self.retrieve(thread_id, run_id) - if run.status in ("cancelled", "cancelling", "completed", "failed", "expired"): - return run - run.tried_cancelling_at = datetime.now(tz=timezone.utc) - self.save(run) - return run + with self._session_factory() as session: + db_run = ( + session.query(RunModel) + .filter( + RunModel.id == run_id, + RunModel.thread_id == thread_id, + ) + .first() + ) + if not db_run: + error_msg = f"Run {run_id} not found in thread {thread_id}" + raise NotFoundError(error_msg) + + run = self._to_pydantic(db_run) + if run.status in ( + "cancelled", + "cancelling", + "completed", + "failed", + "expired", + ): + return run + + db_run.tried_cancelling_at = datetime.now(tz=timezone.utc) + session.commit() + session.refresh(db_run) + return self._to_pydantic(db_run) @override def save(self, run: Run, new: bool = False) -> None: - runs_dir = self.get_runs_dir(run.thread_id) - runs_dir.mkdir(parents=True, exist_ok=True) - run_file = self._get_run_path(run.thread_id, run.id, new=new) - run_file.write_text(run.model_dump_json(), encoding="utf-8") + with self._session_factory() as session: + if new: + # Check if run already exists + existing_run = ( + session.query(RunModel) + .filter( + RunModel.id == run.id, + RunModel.thread_id == run.thread_id, + ) + .first() + ) + if existing_run: + error_msg = f"Run {run.id} already exists in thread {run.thread_id}" + raise ConflictError(error_msg) + + db_run = RunModel.from_pydantic(run) + session.add(db_run) + else: + db_run = ( + session.query(RunModel) + .filter( + RunModel.id == run.id, + RunModel.thread_id == run.thread_id, + ) + .first() + ) + if not db_run: + error_msg = f"Run {run.id} not found in thread {run.thread_id}" + raise NotFoundError(error_msg) + + # Update fields + db_run.status = run.status + db_run.instructions = run.instructions + db_run.tools = run.tools + db_run.metadata = run.metadata + db_run.tried_cancelling_at = run.tried_cancelling_at + db_run.started_at = run.started_at + db_run.completed_at = run.completed_at + db_run.failed_at = run.failed_at + db_run.expired_at = run.expired_at + db_run.cancelled_at = run.cancelled_at + db_run.last_error = run.last_error + db_run.usage = run.usage + + session.commit() diff --git a/src/askui/chat/api/settings.py b/src/askui/chat/api/settings.py index e6ae157a..033598d3 100644 --- a/src/askui/chat/api/settings.py +++ b/src/askui/chat/api/settings.py @@ -1,13 +1,21 @@ from pathlib import Path -from fastmcp.mcp_config import RemoteMCPServer, StdioMCPServer -from pydantic import Field -from pydantic_settings import BaseSettings, SettingsConfigDict - from askui.chat.api.mcp_configs.models import McpConfig from askui.chat.api.telemetry.integrations.fastapi.settings import TelemetrySettings from askui.chat.api.telemetry.logs.settings import LogFilter, LogSettings from askui.utils.datetime_utils import now +from fastmcp.mcp_config import RemoteMCPServer, StdioMCPServer +from pydantic import BaseModel, Field +from pydantic_settings import BaseSettings, SettingsConfigDict + + +class DbSettings(BaseModel): + """Database configuration settings.""" + + url: str = Field( + default_factory=lambda: f"sqlite:///{Path.cwd().absolute()}/askui_chat.db", + description="Database URL for SQLAlchemy connection", + ) def _get_default_mcp_configs(chat_api_host: str, chat_api_port: int) -> list[McpConfig]: @@ -45,8 +53,9 @@ class Settings(BaseSettings): data_dir: Path = Field( default_factory=lambda: Path.cwd() / "chat", - description="Base directory for storing chat data", + description="Base directory for chat data (used during migration)", ) + db: DbSettings = Field(default_factory=DbSettings) host: str = Field( default="127.0.0.1", description="Host for the chat API", diff --git a/src/askui/chat/api/threads/dependencies.py b/src/askui/chat/api/threads/dependencies.py index 64ff8172..d8d33518 100644 --- a/src/askui/chat/api/threads/dependencies.py +++ b/src/askui/chat/api/threads/dependencies.py @@ -1,24 +1,21 @@ -from pathlib import Path - -from fastapi import Depends - -from askui.chat.api.dependencies import WorkspaceDirDep +from askui.chat.api.dependencies import SessionFactoryDep from askui.chat.api.messages.dependencies import MessageServiceDep from askui.chat.api.messages.service import MessageService from askui.chat.api.runs.dependencies import RunServiceDep from askui.chat.api.runs.service import RunService from askui.chat.api.threads.facade import ThreadFacade from askui.chat.api.threads.service import ThreadService +from fastapi import Depends def get_thread_service( - workspace_dir: Path = WorkspaceDirDep, + session_factory=SessionFactoryDep, message_service: MessageService = MessageServiceDep, run_service: RunService = RunServiceDep, ) -> ThreadService: """Get ThreadService instance.""" return ThreadService( - base_dir=workspace_dir, + session_factory=session_factory, message_service=message_service, run_service=run_service, ) diff --git a/src/askui/chat/api/threads/facade.py b/src/askui/chat/api/threads/facade.py index de836dd4..e0e05f35 100644 --- a/src/askui/chat/api/threads/facade.py +++ b/src/askui/chat/api/threads/facade.py @@ -1,10 +1,10 @@ from collections.abc import AsyncGenerator -from askui.chat.api.messages.models import Message, MessageCreateParams +from askui.chat.api.messages.schemas import Message, MessageCreateParams from askui.chat.api.messages.service import MessageService from askui.chat.api.models import ThreadId, WorkspaceId from askui.chat.api.runs.events.events import Event -from askui.chat.api.runs.models import ( +from askui.chat.api.runs.schemas import ( Run, RunCreateParams, RunListQuery, diff --git a/src/askui/chat/api/threads/models.py b/src/askui/chat/api/threads/models.py index 6ee1931f..db0d272d 100644 --- a/src/askui/chat/api/threads/models.py +++ b/src/askui/chat/api/threads/models.py @@ -1,52 +1,54 @@ -from typing import Literal +"""Thread database model.""" -from pydantic import BaseModel +from datetime import datetime, timezone -from askui.chat.api.messages.models import MessageCreateParams -from askui.chat.api.models import ThreadId -from askui.utils.api_utils import Resource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import generate_time_ordered_id -from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import ThreadId +from askui.chat.api.threads.schemas import Thread, ThreadCreateParams +from bson import ObjectId +from sqlalchemy import Column, DateTime, String -class ThreadBase(BaseModel): - """Base thread model.""" +class ThreadModel(Base): + """Thread database model.""" - name: str | None = None + __tablename__ = "threads" + id = Column(ThreadId, primary_key=True) + created_at = Column(DateTime, nullable=False, index=True) + name = Column(String(128), nullable=True) + @staticmethod + def create_id() -> str: + """Create a new thread ID with prefix.""" + return f"thread_{ObjectId()}" -class ThreadCreateParams(ThreadBase): - """Parameters for creating a thread.""" + def to_pydantic(self) -> Thread: + """Convert to Pydantic model.""" + # Ensure created_at is timezone-aware + created_at = self.created_at + if created_at.tzinfo is None: + created_at = created_at.replace(tzinfo=timezone.utc) - messages: list[MessageCreateParams] | None = None - - -class ThreadModifyParams(BaseModelWithNotGiven): - """Parameters for modifying a thread.""" - - name: str | None | NotGiven = NOT_GIVEN - - -class Thread(ThreadBase, Resource): - """A chat thread/session.""" - - id: ThreadId - object: Literal["thread"] = "thread" - created_at: UnixDatetime + return Thread( + id=self.id, # Prefix is handled by the specialized type + created_at=created_at, + name=self.name, + ) @classmethod - def create(cls, params: ThreadCreateParams) -> "Thread": + def from_pydantic(cls, thread: Thread) -> "ThreadModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("thread"), - created_at=now(), - **params.model_dump(exclude={"messages"}), + id=thread.id, + created_at=thread.created_at, + name=thread.name, ) - def modify(self, params: ThreadModifyParams) -> "Thread": - return Thread.model_validate( - { - **self.model_dump(), - **params.model_dump(), - } + @classmethod + def from_create_params(cls, params: ThreadCreateParams) -> "ThreadModel": + """Create from create parameters.""" + return cls( + id=cls.create_id(), + created_at=datetime.now(timezone.utc), + name=params.name, ) diff --git a/src/askui/chat/api/threads/router.py b/src/askui/chat/api/threads/router.py index a9e18bf4..472be899 100644 --- a/src/askui/chat/api/threads/router.py +++ b/src/askui/chat/api/threads/router.py @@ -1,11 +1,14 @@ -from fastapi import APIRouter, status - from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.models import ThreadId from askui.chat.api.threads.dependencies import ThreadServiceDep -from askui.chat.api.threads.models import Thread, ThreadCreateParams, ThreadModifyParams +from askui.chat.api.threads.schemas import ( + Thread, + ThreadCreateParams, + ThreadModifyParams, +) from askui.chat.api.threads.service import ThreadService from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, status router = APIRouter(prefix="/threads", tags=["threads"]) diff --git a/src/askui/chat/api/threads/schemas.py b/src/askui/chat/api/threads/schemas.py new file mode 100644 index 00000000..cd194af4 --- /dev/null +++ b/src/askui/chat/api/threads/schemas.py @@ -0,0 +1,51 @@ +from typing import Literal + +from askui.chat.api.messages.schemas import MessageCreateParams +from askui.chat.api.models import ThreadId +from askui.utils.api_utils import Resource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import generate_time_ordered_id +from askui.utils.not_given import NOT_GIVEN, BaseModelWithNotGiven, NotGiven +from pydantic import BaseModel + + +class ThreadBase(BaseModel): + """Base thread model.""" + + name: str | None = None + + +class ThreadCreateParams(ThreadBase): + """Parameters for creating a thread.""" + + messages: list[MessageCreateParams] | None = None + + +class ThreadModifyParams(BaseModelWithNotGiven): + """Parameters for modifying a thread.""" + + name: str | None | NotGiven = NOT_GIVEN + + +class Thread(ThreadBase, Resource): + """A chat thread/session.""" + + id: ThreadId + object: Literal["thread"] = "thread" + created_at: UnixDatetime + + @classmethod + def create(cls, params: ThreadCreateParams) -> "Thread": + return cls( + id=generate_time_ordered_id("thread"), + created_at=now(), + **params.model_dump(exclude={"messages"}), + ) + + def modify(self, params: ThreadModifyParams) -> "Thread": + return Thread.model_validate( + { + **self.model_dump(), + **params.model_dump(), + } + ) diff --git a/src/askui/chat/api/threads/service.py b/src/askui/chat/api/threads/service.py index ca54f203..2371f5e0 100644 --- a/src/askui/chat/api/threads/service.py +++ b/src/askui/chat/api/threads/service.py @@ -1,82 +1,107 @@ -import shutil -from pathlib import Path +from typing import Callable +from askui.chat.api.db.query_builder import QueryBuilder from askui.chat.api.messages.service import MessageService from askui.chat.api.models import ThreadId from askui.chat.api.runs.service import RunService -from askui.chat.api.threads.models import Thread, ThreadCreateParams, ThreadModifyParams -from askui.utils.api_utils import ( - ConflictError, - ListQuery, - ListResponse, - NotFoundError, - list_resources, +from askui.chat.api.threads.models import ThreadModel +from askui.chat.api.threads.schemas import ( + Thread, + ThreadCreateParams, + ThreadModifyParams, ) +from askui.utils.api_utils import ListQuery, ListResponse, NotFoundError +from sqlalchemy.orm import Session class ThreadService: - """Service for managing Thread resources with filesystem persistence.""" + """Service for managing Thread resources with SQLAlchemy persistence.""" def __init__( - self, base_dir: Path, message_service: MessageService, run_service: RunService + self, + session_factory: Callable[[], Session], + message_service: MessageService, + run_service: RunService, ) -> None: - self._base_dir = base_dir - self._threads_dir = base_dir / "threads" + self._session_factory = session_factory self._message_service = message_service self._run_service = run_service - def _get_thread_path(self, thread_id: ThreadId, new: bool = False) -> Path: - thread_path = self._threads_dir / f"{thread_id}.json" - exists = thread_path.exists() - if new and exists: - error_msg = f"Thread {thread_id} already exists" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"Thread {thread_id} not found" - raise NotFoundError(error_msg) - return thread_path + def _to_pydantic(self, db_model: ThreadModel) -> Thread: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_(self, query: ListQuery) -> ListResponse[Thread]: - return list_resources(self._threads_dir, query, Thread) + with self._session_factory() as session: + q = session.query(ThreadModel) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, ThreadModel, query, ThreadModel.created_at, ThreadModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve(self, thread_id: ThreadId) -> Thread: - try: - thread_path = self._get_thread_path(thread_id) - return Thread.model_validate_json(thread_path.read_text()) - except FileNotFoundError as e: - error_msg = f"Thread {thread_id} not found" - raise NotFoundError(error_msg) from e + with self._session_factory() as session: + db_thread = ( + session.query(ThreadModel).filter(ThreadModel.id == thread_id).first() + ) + if not db_thread: + error_msg = f"Thread {thread_id} not found" + raise NotFoundError(error_msg) + return self._to_pydantic(db_thread) def create(self, params: ThreadCreateParams) -> Thread: - thread = Thread.create(params) - self._save(thread, new=True) + with self._session_factory() as session: + db_thread = ThreadModel.from_create_params(params) + session.add(db_thread) + session.commit() + session.refresh(db_thread) + + thread = self._to_pydantic(db_thread) - if params.messages: - for message in params.messages: - self._message_service.create( - thread_id=thread.id, - params=message, - ) - return thread + if params.messages: + for message in params.messages: + self._message_service.create( + thread_id=thread.id, + params=message, + ) + return thread def modify(self, thread_id: ThreadId, params: ThreadModifyParams) -> Thread: - thread = self.retrieve(thread_id) - modified = thread.modify(params) - self._save(modified) - return modified + with self._session_factory() as session: + db_thread = ( + session.query(ThreadModel).filter(ThreadModel.id == thread_id).first() + ) + if not db_thread: + error_msg = f"Thread {thread_id} not found" + raise NotFoundError(error_msg) + + # Update fields + if params.name is not None: + db_thread.name = params.name + + session.commit() + session.refresh(db_thread) + + return self._to_pydantic(db_thread) def delete(self, thread_id: ThreadId) -> None: - try: - shutil.rmtree( - self._message_service.get_messages_dir(thread_id), ignore_errors=True + with self._session_factory() as session: + db_thread = ( + session.query(ThreadModel).filter(ThreadModel.id == thread_id).first() ) - shutil.rmtree(self._run_service.get_runs_dir(thread_id), ignore_errors=True) - self._get_thread_path(thread_id).unlink() - except FileNotFoundError as e: - error_msg = f"Thread {thread_id} not found" - raise NotFoundError(error_msg) from e - - def _save(self, thread: Thread, new: bool = False) -> None: - self._threads_dir.mkdir(parents=True, exist_ok=True) - thread_file = self._get_thread_path(thread.id, new=new) - thread_file.write_text(thread.model_dump_json(), encoding="utf-8") + if not db_thread: + error_msg = f"Thread {thread_id} not found" + raise NotFoundError(error_msg) + + # Delete related messages and runs (cascade will handle this) + session.delete(db_thread) + session.commit() + session.commit() diff --git a/src/askui/chat/api/utils.py b/src/askui/chat/api/utils.py index caf1586a..595099f7 100644 --- a/src/askui/chat/api/utils.py +++ b/src/askui/chat/api/utils.py @@ -11,3 +11,28 @@ def filter_fn(resource: WorkspaceResourceT) -> bool: return resource.workspace_id is None or resource.workspace_id == workspace return filter_fn + + +def add_prefix(prefix: str, object_id: str) -> str: + """Add prefix to ObjectId. + + Args: + prefix (str): Prefix to add (e.g., "asst", "thread"). + object_id (str): ObjectId without prefix. + + Returns: + str: Prefixed ID. + """ + return f"{prefix}_{object_id}" + + +def remove_prefix(prefixed_id: str) -> str: + """Remove prefix from ObjectId. + + Args: + prefixed_id (str): Prefixed ID (e.g., "asst_507f1f77bcf86cd799439011"). + + Returns: + str: ObjectId without prefix. + """ + return prefixed_id.split("_", 1)[1] diff --git a/src/askui/chat/api/workflows/models.py b/src/askui/chat/api/workflows/models.py index daf984bb..20770993 100644 --- a/src/askui/chat/api/workflows/models.py +++ b/src/askui/chat/api/workflows/models.py @@ -1,73 +1,65 @@ -from typing import Annotated, Literal +"""Workflow database models.""" -from pydantic import BaseModel, Field +from askui.chat.api.db.base import Base +from askui.chat.api.db.types import WorkflowId +from askui.chat.api.workflows.schemas import Workflow +from sqlalchemy import Column, DateTime, ForeignKey, Index, Integer, String +from sqlalchemy.orm import relationship -from askui.chat.api.models import WorkspaceId, WorkspaceResource -from askui.utils.datetime_utils import UnixDatetime, now -from askui.utils.id_utils import IdField, generate_time_ordered_id -WorkflowId = Annotated[str, IdField("wf")] +class WorkflowModel(Base): + """Workflow database model.""" + __tablename__ = "workflows" + id = Column(WorkflowId, primary_key=True) + workspace_id = Column(String(36), nullable=True, index=True) + created_at = Column(DateTime, nullable=False, index=True) + name = Column(String, nullable=False) + description = Column(String, nullable=False) + # Relationship to tags + tags = relationship( + "WorkflowTagModel", + back_populates="workflow", + cascade="all, delete-orphan", + lazy="joined", + ) -class WorkflowCreateParams(BaseModel): - """ - Parameters for creating a workflow via API. - """ - - name: str - description: str - tags: list[str] = Field(default_factory=list) - - -class WorkflowModifyParams(BaseModel): - """ - Parameters for modifying a workflow via API. - """ - - name: str | None = None - description: str | None = None - tags: list[str] | None = None - - -class Workflow(WorkspaceResource): - """ - A workflow resource in the chat API. - - Args: - id (WorkflowId): The id of the workflow. Must start with the 'wf_' prefix and be - followed by one or more alphanumerical characters. - object (Literal['workflow']): The object type, always 'workflow'. - created_at (UnixDatetime): The creation time as a Unix timestamp. - name (str): The name or title of the workflow. - description (str): A detailed description of the workflow's purpose and steps. - tags (list[str], optional): Tags associated with the workflow for filtering or - categorization. Default is an empty list. - workspace_id (WorkspaceId | None, optional): The workspace this workflow belongs to. - """ - - id: WorkflowId - object: Literal["workflow"] = "workflow" - created_at: UnixDatetime - name: str - description: str - tags: list[str] = Field(default_factory=list) + def to_pydantic(self) -> Workflow: + """Convert to Pydantic model.""" + data = { + "id": self.id, # Prefix is handled by the specialized type + "workspace_id": self.workspace_id, + "created_at": self.created_at, + "name": self.name, + "description": self.description, + "tags": [tag.tag for tag in self.tags], + } + return Workflow.model_validate(data) @classmethod - def create( - cls, workspace_id: WorkspaceId | None, params: WorkflowCreateParams - ) -> "Workflow": + def from_pydantic(cls, workflow: Workflow) -> "WorkflowModel": + """Create from Pydantic model.""" return cls( - id=generate_time_ordered_id("wf"), - created_at=now(), - workspace_id=workspace_id, - **params.model_dump(), + id=workflow.id, + workspace_id=str(workflow.workspace_id) if workflow.workspace_id else None, + created_at=workflow.created_at, + name=workflow.name, + description=workflow.description, ) - def modify(self, params: WorkflowModifyParams) -> "Workflow": - update_data = {k: v for k, v in params.model_dump().items() if v is not None} - return Workflow.model_validate( - { - **self.model_dump(), - **update_data, - } - ) + +class WorkflowTagModel(Base): + """Workflow tag database model.""" + + __tablename__ = "workflow_tags" + id = Column(Integer, primary_key=True, autoincrement=True) + workflow_id = Column( + WorkflowId, + ForeignKey("workflows.id", ondelete="CASCADE"), + nullable=False, + index=True, + ) + tag = Column(String, nullable=False, index=True) + workflow = relationship("WorkflowModel", back_populates="tags") + + __table_args__ = (Index("idx_workflow_tags_tag_workflow", "tag", "workflow_id"),) diff --git a/src/askui/chat/api/workflows/router.py b/src/askui/chat/api/workflows/router.py index c0546db7..69695dcc 100644 --- a/src/askui/chat/api/workflows/router.py +++ b/src/askui/chat/api/workflows/router.py @@ -1,11 +1,9 @@ from typing import Annotated -from fastapi import APIRouter, Header, Path, Query, status - from askui.chat.api.dependencies import ListQueryDep from askui.chat.api.models import WorkspaceId from askui.chat.api.workflows.dependencies import WorkflowServiceDep -from askui.chat.api.workflows.models import ( +from askui.chat.api.workflows.schemas import ( Workflow, WorkflowCreateParams, WorkflowId, @@ -13,6 +11,7 @@ ) from askui.chat.api.workflows.service import WorkflowService from askui.utils.api_utils import ListQuery, ListResponse +from fastapi import APIRouter, Header, Path, Query, status router = APIRouter(prefix="/workflows", tags=["workflows"]) diff --git a/src/askui/chat/api/workflows/schemas.py b/src/askui/chat/api/workflows/schemas.py new file mode 100644 index 00000000..daf984bb --- /dev/null +++ b/src/askui/chat/api/workflows/schemas.py @@ -0,0 +1,73 @@ +from typing import Annotated, Literal + +from pydantic import BaseModel, Field + +from askui.chat.api.models import WorkspaceId, WorkspaceResource +from askui.utils.datetime_utils import UnixDatetime, now +from askui.utils.id_utils import IdField, generate_time_ordered_id + +WorkflowId = Annotated[str, IdField("wf")] + + +class WorkflowCreateParams(BaseModel): + """ + Parameters for creating a workflow via API. + """ + + name: str + description: str + tags: list[str] = Field(default_factory=list) + + +class WorkflowModifyParams(BaseModel): + """ + Parameters for modifying a workflow via API. + """ + + name: str | None = None + description: str | None = None + tags: list[str] | None = None + + +class Workflow(WorkspaceResource): + """ + A workflow resource in the chat API. + + Args: + id (WorkflowId): The id of the workflow. Must start with the 'wf_' prefix and be + followed by one or more alphanumerical characters. + object (Literal['workflow']): The object type, always 'workflow'. + created_at (UnixDatetime): The creation time as a Unix timestamp. + name (str): The name or title of the workflow. + description (str): A detailed description of the workflow's purpose and steps. + tags (list[str], optional): Tags associated with the workflow for filtering or + categorization. Default is an empty list. + workspace_id (WorkspaceId | None, optional): The workspace this workflow belongs to. + """ + + id: WorkflowId + object: Literal["workflow"] = "workflow" + created_at: UnixDatetime + name: str + description: str + tags: list[str] = Field(default_factory=list) + + @classmethod + def create( + cls, workspace_id: WorkspaceId | None, params: WorkflowCreateParams + ) -> "Workflow": + return cls( + id=generate_time_ordered_id("wf"), + created_at=now(), + workspace_id=workspace_id, + **params.model_dump(), + ) + + def modify(self, params: WorkflowModifyParams) -> "Workflow": + update_data = {k: v for k, v in params.model_dump().items() if v is not None} + return Workflow.model_validate( + { + **self.model_dump(), + **update_data, + } + ) diff --git a/src/askui/chat/api/workflows/service.py b/src/askui/chat/api/workflows/service.py index 29aaa1aa..1adb6de1 100644 --- a/src/askui/chat/api/workflows/service.py +++ b/src/askui/chat/api/workflows/service.py @@ -1,56 +1,27 @@ -from pathlib import Path from typing import Callable +from askui.chat.api.db.query_builder import QueryBuilder from askui.chat.api.models import WorkspaceId -from askui.chat.api.utils import build_workspace_filter_fn -from askui.chat.api.workflows.models import ( +from askui.chat.api.workflows.models import WorkflowModel, WorkflowTagModel +from askui.chat.api.workflows.schemas import ( Workflow, WorkflowCreateParams, WorkflowId, WorkflowModifyParams, ) -from askui.utils.api_utils import ( - ConflictError, - ListQuery, - ListResponse, - NotFoundError, - list_resources, -) - - -def _build_workflow_filter_fn( - workspace_id: WorkspaceId | None, - tags: list[str] | None = None, -) -> Callable[[Workflow], bool]: - workspace_filter: Callable[[Workflow], bool] = build_workspace_filter_fn( - workspace_id, Workflow - ) +from askui.utils.api_utils import ForbiddenError, ListQuery, ListResponse, NotFoundError +from sqlalchemy.orm import Session - def filter_fn(workflow: Workflow) -> bool: - if not workspace_filter(workflow): - return False - if tags is not None: - return any(tag in workflow.tags for tag in tags) - return True - return filter_fn +class WorkflowService: + """Service for managing Workflow resources with SQLAlchemy persistence.""" + def __init__(self, session_factory: Callable[[], Session]) -> None: + self._session_factory = session_factory -class WorkflowService: - def __init__(self, base_dir: Path) -> None: - self._base_dir = base_dir - self._workflows_dir = base_dir / "workflows" - - def _get_workflow_path(self, workflow_id: WorkflowId, new: bool = False) -> Path: - workflow_path = self._workflows_dir / f"{workflow_id}.json" - exists = workflow_path.exists() - if new and exists: - error_msg = f"Workflow {workflow_id} already exists" - raise ConflictError(error_msg) - if not new and not exists: - error_msg = f"Workflow {workflow_id} not found" - raise NotFoundError(error_msg) - return workflow_path + def _to_pydantic(self, db_model: WorkflowModel) -> Workflow: + """Convert SQLAlchemy model to Pydantic model.""" + return db_model.to_pydantic() def list_( self, @@ -58,37 +29,74 @@ def list_( query: ListQuery, tags: list[str] | None = None, ) -> ListResponse[Workflow]: - return list_resources( - base_dir=self._workflows_dir, - query=query, - resource_type=Workflow, - filter_fn=_build_workflow_filter_fn(workspace_id, tags=tags), - ) + with self._session_factory() as session: + q = session.query(WorkflowModel) + + # Filter by workspace + if workspace_id is not None: + q = q.filter(WorkflowModel.workspace_id == str(workspace_id)) + else: + q = q.filter(WorkflowModel.workspace_id.is_(None)) + + # Filter by tags if specified + if tags: + q = q.join(WorkflowTagModel).filter(WorkflowTagModel.tag.in_(tags)) + + # Apply list query parameters + q = QueryBuilder.apply_list_query( + q, WorkflowModel, query, WorkflowModel.created_at, WorkflowModel.id + ) + + # Apply limit + limit = query.limit or 20 + q = q.limit(limit + 1) # +1 to check if there are more + + results = q.all() + return QueryBuilder.build_list_response(results, limit, self._to_pydantic) def retrieve( self, workspace_id: WorkspaceId | None, workflow_id: WorkflowId ) -> Workflow: - try: - workflow_path = self._get_workflow_path(workflow_id) - workflow = Workflow.model_validate_json(workflow_path.read_text()) + with self._session_factory() as session: + db_workflow = ( + session.query(WorkflowModel) + .filter(WorkflowModel.id == workflow_id) + .first() + ) + if not db_workflow: + error_msg = f"Workflow {workflow_id} not found" + raise NotFoundError(error_msg) # Check workspace access - if workspace_id is not None and workflow.workspace_id != workspace_id: + if not ( + db_workflow.workspace_id is None + or db_workflow.workspace_id == str(workspace_id) + ): error_msg = f"Workflow {workflow_id} not found" raise NotFoundError(error_msg) - except FileNotFoundError as e: - error_msg = f"Workflow {workflow_id} not found" - raise NotFoundError(error_msg) from e - else: - return workflow + return self._to_pydantic(db_workflow) def create( self, workspace_id: WorkspaceId | None, params: WorkflowCreateParams ) -> Workflow: - workflow = Workflow.create(workspace_id, params) - self._save(workflow, new=True) - return workflow + with self._session_factory() as session: + db_workflow = WorkflowModel.from_create_params(params, workspace_id) + session.add(db_workflow) + session.commit() + session.refresh(db_workflow) + + # Add tags + if params.tags: + for tag in params.tags: + db_tag = WorkflowTagModel( + workflow_id=db_workflow.id, + tag=tag, + ) + session.add(db_tag) + + session.commit() + return self._to_pydantic(db_workflow) def modify( self, @@ -96,12 +104,88 @@ def modify( workflow_id: WorkflowId, params: WorkflowModifyParams, ) -> Workflow: - workflow = self.retrieve(workspace_id, workflow_id) - modified = workflow.modify(params) - self._save(modified) - return modified - - def _save(self, workflow: Workflow, new: bool = False) -> None: - self._workflows_dir.mkdir(parents=True, exist_ok=True) - workflow_file = self._get_workflow_path(workflow.id, new=new) - workflow_file.write_text(workflow.model_dump_json(), encoding="utf-8") + with self._session_factory() as session: + db_workflow = ( + session.query(WorkflowModel) + .filter(WorkflowModel.id == workflow_id) + .first() + ) + if not db_workflow: + error_msg = f"Workflow {workflow_id} not found" + raise NotFoundError(error_msg) + + # Check workspace access + if not ( + db_workflow.workspace_id is None + or db_workflow.workspace_id == str(workspace_id) + ): + error_msg = f"Workflow {workflow_id} not found" + raise NotFoundError(error_msg) + + if db_workflow.workspace_id is None: + error_msg = f"Default workflow {workflow_id} cannot be modified" + raise ForbiddenError(error_msg) + + # Update fields + if params.name is not None: + db_workflow.name = params.name + if params.description is not None: + db_workflow.description = params.description + if params.steps is not None: + db_workflow.steps = params.steps + + # Update tags if provided + if params.tags is not None: + # Remove existing tags + session.query(WorkflowTagModel).filter( + WorkflowTagModel.workflow_id == workflow_id + ).delete() + + # Add new tags + for tag in params.tags: + db_tag = WorkflowTagModel( + workflow_id=workflow_id, + tag=tag, + ) + session.add(db_tag) + + session.commit() + session.refresh(db_workflow) + return self._to_pydantic(db_workflow) + + def delete( + self, + workspace_id: WorkspaceId | None, + workflow_id: WorkflowId, + force: bool = False, + ) -> None: + with self._session_factory() as session: + db_workflow = ( + session.query(WorkflowModel) + .filter(WorkflowModel.id == workflow_id) + .first() + ) + if not db_workflow: + error_msg = f"Workflow {workflow_id} not found" + if not force: + raise NotFoundError(error_msg) + return + + # Check workspace access + if not ( + db_workflow.workspace_id is None + or db_workflow.workspace_id == str(workspace_id) + ): + error_msg = f"Workflow {workflow_id} not found" + if not force: + raise NotFoundError(error_msg) + return + + if db_workflow.workspace_id is None and not force: + error_msg = f"Default workflow {workflow_id} cannot be deleted" + raise ForbiddenError(error_msg) + + # Delete related tags (cascade will handle this) + session.delete(db_workflow) + session.commit() + session.commit() diff --git a/src/askui/chat/migrations/__init__.py b/src/askui/chat/migrations/__init__.py new file mode 100644 index 00000000..57b6e4b1 --- /dev/null +++ b/src/askui/chat/migrations/__init__.py @@ -0,0 +1,134 @@ +"""Migration framework for chat API.""" + +from datetime import datetime, timezone +from pathlib import Path +from typing import Callable + +from askui.chat.api.db.base import Base +from askui.chat.api.migrations.models import MigrationVersionModel +from sqlalchemy import create_engine, text +from sqlalchemy.orm import sessionmaker + + +class MigrationRunner: + """Handles database migrations for chat API. + + Uses version-based migrations stored in the database rather than + external migration files like Alembic. + """ + + def __init__(self, database_url: str) -> None: + """Initialize migration runner. + + Args: + database_url (str): Database URL for SQLAlchemy connection. + """ + self.database_url = database_url + self.engine = create_engine(database_url) + + def get_current_version(self) -> int: + """Get current schema version from database. + + Returns: + int: Current schema version, or 0 if no version found. + """ + with self.engine.connect() as conn: + try: + result = conn.execute( + text("SELECT MAX(version) FROM migration_version") + ) + version = result.scalar() + return version if version else 0 + except Exception: + return 0 + + def should_migrate(self, data_dir: Path) -> bool: + """Check if migration is needed. + + Args: + data_dir (Path): Directory containing JSON/JSONL files. + + Returns: + bool: True if migration is needed. + """ + current_version = self.get_current_version() + + # Import migrations here to avoid circular imports + from .versions import MIGRATIONS + + target_version = max(MIGRATIONS.keys()) + + # Need to migrate if version is behind + if current_version < target_version: + return True + + # Also check if JSON files exist (for initial migration) + if current_version == 0: + json_dirs = [ + "assistants", + "threads", + "messages", + "runs", + "files", + "mcp_configs", + "workflows", + ] + for dir_name in json_dirs: + dir_path = data_dir / dir_name + if dir_path.exists(): + json_files = list(dir_path.glob("*.json")) + if json_files: + return True + + return False + + def migrate(self, data_dir: Path) -> None: + """Run all pending migrations. + + Args: + data_dir (Path): Directory containing JSON/JSONL files. + """ + current_version = self.get_current_version() + + # Import migrations here to avoid circular imports + from .versions import MIGRATIONS + + target_version = max(MIGRATIONS.keys()) + + for version in range(current_version + 1, target_version + 1): + if version in MIGRATIONS: + migration = MIGRATIONS[version] + print(f"Running migration {version}: {migration.__name__}") + migration.upgrade(self.engine, data_dir) + + # Record version + with self.engine.begin() as conn: + conn.execute( + text( + "INSERT INTO migration_version (version, applied_at) VALUES (:v, :t)" + ), + {"v": version, "t": datetime.now(timezone.utc)}, + ) + + print("Migration completed successfully") + + +# Migration functions registry +MIGRATIONS: dict[int, Callable] = {} + + +def register_migration(version: int) -> Callable: + """Decorator to register a migration function. + + Args: + version (int): Migration version number. + + Returns: + Callable: Decorator function. + """ + + def decorator(func: Callable) -> Callable: + MIGRATIONS[version] = func + return func + + return decorator diff --git a/src/askui/chat/migrations/versions/__init__.py b/src/askui/chat/migrations/versions/__init__.py new file mode 100644 index 00000000..3a058bd6 --- /dev/null +++ b/src/askui/chat/migrations/versions/__init__.py @@ -0,0 +1,8 @@ +"""Migration versions registry.""" + +from . import migration_001_initial_schema, migration_002_migrate_json_to_sqlite + +MIGRATIONS = { + 1: migration_001_initial_schema, + 2: migration_002_migrate_json_to_sqlite, +} diff --git a/src/askui/chat/migrations/versions/migration_001_initial_schema.py b/src/askui/chat/migrations/versions/migration_001_initial_schema.py new file mode 100644 index 00000000..1a1ddf98 --- /dev/null +++ b/src/askui/chat/migrations/versions/migration_001_initial_schema.py @@ -0,0 +1,11 @@ +"""Initial schema migration.""" + +from askui.chat.migrations import register_migration + + +@register_migration(1) +def upgrade(engine, data_dir): + """Create all tables.""" + from askui.chat.api.db.base import Base + + Base.metadata.create_all(engine) diff --git a/src/askui/chat/migrations/versions/migration_002_migrate_json_to_sqlite.py b/src/askui/chat/migrations/versions/migration_002_migrate_json_to_sqlite.py new file mode 100644 index 00000000..f5bc0bb1 --- /dev/null +++ b/src/askui/chat/migrations/versions/migration_002_migrate_json_to_sqlite.py @@ -0,0 +1,294 @@ +"""Data migration from JSON/JSONL to SQLite.""" + +import json +import logging +from datetime import datetime, timezone +from pathlib import Path + +from askui.chat.api.assistants.models import AssistantModel +from askui.chat.api.assistants.schemas import Assistant +from askui.chat.api.files.models import FileModel +from askui.chat.api.files.schemas import File as FilePydantic +from askui.chat.api.mcp_configs.models import McpConfigModel +from askui.chat.api.mcp_configs.schemas import McpConfig +from askui.chat.api.messages.models import MessageModel +from askui.chat.api.messages.schemas import Message +from askui.chat.api.runs.events.events import Event +from askui.chat.api.runs.events.models import EventModel +from askui.chat.api.runs.models import RunModel +from askui.chat.api.runs.schemas import Run +from askui.chat.api.threads.models import ThreadModel +from askui.chat.api.threads.schemas import Thread +from askui.chat.api.workflows.models import WorkflowModel, WorkflowTagModel +from askui.chat.api.workflows.schemas import Workflow +from askui.utils.datetime_utils import now +from sqlalchemy.orm import sessionmaker + +logger = logging.getLogger(__name__) + + +def upgrade(engine, data_dir: Path): + """Migrate data from JSON/JSONL to SQLite.""" + Session = sessionmaker(bind=engine) + session = Session() + + try: + # Migrate assistants + assistants_dir = data_dir / "assistants" + if assistants_dir.exists(): + for json_file in assistants_dir.glob("*.json"): + try: + assistant = Assistant.model_validate_json(json_file.read_text()) + # Extract ObjectId from prefixed ID + object_id = assistant.id.split("_", 1)[1] + + # Check if assistant already exists + existing = ( + session.query(AssistantModel) + .filter(AssistantModel.id == object_id) + .first() + ) + if existing: + logger.info(f"Assistant {object_id} already exists, skipping") + continue + + db_assistant = AssistantModel( + id=object_id, + workspace_id=str(assistant.workspace_id) + if assistant.workspace_id + else None, + created_at=assistant.created_at, + name=assistant.name, + description=assistant.description, + avatar=assistant.avatar, + tools=assistant.tools, + system=assistant.system, + ) + session.add(db_assistant) + except Exception as e: + logger.warning(f"Failed to migrate assistant {json_file}: {e}") + + # Migrate threads + threads_dir = data_dir / "threads" + if threads_dir.exists(): + for json_file in threads_dir.glob("*.json"): + try: + thread = Thread.model_validate_json(json_file.read_text()) + object_id = thread.id.split("_", 1)[1] + + # Check if thread already exists + existing = ( + session.query(ThreadModel) + .filter(ThreadModel.id == object_id) + .first() + ) + if existing: + logger.info(f"Thread {object_id} already exists, skipping") + continue + + db_thread = ThreadModel( + id=object_id, + created_at=thread.created_at, + name=thread.name, + ) + session.add(db_thread) + except Exception as e: + logger.warning(f"Failed to migrate thread {json_file}: {e}") + + # Migrate files + files_dir = data_dir / "files" + if files_dir.exists(): + for json_file in files_dir.glob("*.json"): + try: + file_data = json.loads(json_file.read_text()) + # Handle incomplete file data from tests + if "size" not in file_data: + file_data["size"] = 0 + if "media_type" not in file_data: + file_data["media_type"] = "application/octet-stream" + if "filename" not in file_data: + file_data["filename"] = "unknown" + if "created_at" not in file_data: + file_data["created_at"] = datetime.now(timezone.utc).isoformat() + + file_model = FilePydantic.model_validate(file_data) + object_id = file_model.id.split("_", 1)[1] + + # Check if file already exists + existing = ( + session.query(FileModel) + .filter(FileModel.id == object_id) + .first() + ) + if existing: + logger.info(f"File {object_id} already exists, skipping") + continue + + db_file = FileModel( + id=object_id, + created_at=file_model.created_at, + filename=file_model.filename, + size=file_model.size, + media_type=file_model.media_type, + ) + session.add(db_file) + except Exception as e: + logger.warning(f"Failed to migrate file {json_file}: {e}") + + # Migrate MCP configs + mcp_configs_dir = data_dir / "mcp_configs" + if mcp_configs_dir.exists(): + for json_file in mcp_configs_dir.glob("*.json"): + try: + mcp_config = McpConfig.model_validate_json(json_file.read_text()) + object_id = mcp_config.id.split("_", 1)[1] + + # Check if MCP config already exists + existing = ( + session.query(McpConfigModel) + .filter(McpConfigModel.id == object_id) + .first() + ) + if existing: + logger.info(f"MCP config {object_id} already exists, skipping") + continue + + db_mcp_config = McpConfigModel( + id=object_id, + workspace_id=str(mcp_config.workspace_id) + if mcp_config.workspace_id + else None, + created_at=mcp_config.created_at, + name=mcp_config.name, + mcp_server=mcp_config.mcp_server.model_dump(), + ) + session.add(db_mcp_config) + except Exception as e: + logger.warning(f"Failed to migrate MCP config {json_file}: {e}") + + # Migrate workflows with tags + workflows_dir = data_dir / "workflows" + if workflows_dir.exists(): + for json_file in workflows_dir.glob("*.json"): + workflow = Workflow.model_validate_json(json_file.read_text()) + object_id = workflow.id.split("_", 1)[1] + db_workflow = WorkflowModel( + id=object_id, + workspace_id=str(workflow.workspace_id) + if workflow.workspace_id + else None, + created_at=workflow.created_at, + name=workflow.name, + description=workflow.description, + ) + session.add(db_workflow) + + # Add tags + for tag in workflow.tags: + db_tag = WorkflowTagModel(workflow_id=object_id, tag=tag) + session.add(db_tag) + + # Migrate messages + messages_dir = data_dir / "messages" + if messages_dir.exists(): + for thread_dir in messages_dir.iterdir(): + if thread_dir.is_dir(): + thread_id = thread_dir.name.split("_", 1)[1] # Remove prefix + for json_file in thread_dir.glob("*.json"): + message = Message.model_validate_json(json_file.read_text()) + object_id = message.id.split("_", 1)[1] + db_message = MessageModel( + id=object_id, + thread_id=thread_id, + created_at=message.created_at, + assistant_id=message.assistant_id.split("_", 1)[1] + if message.assistant_id + else None, + run_id=message.run_id.split("_", 1)[1] + if message.run_id + else None, + role=message.role, + content=message.content, + stop_reason=message.stop_reason, + ) + session.add(db_message) + + # Migrate runs + runs_dir = data_dir / "runs" + if runs_dir.exists(): + for thread_dir in runs_dir.iterdir(): + if thread_dir.is_dir(): + thread_id = thread_dir.name.split("_", 1)[1] # Remove prefix + for json_file in thread_dir.glob("*.json"): + run = Run.model_validate_json(json_file.read_text()) + object_id = run.id.split("_", 1)[1] + db_run = RunModel( + id=object_id, + thread_id=thread_id, + assistant_id=run.assistant_id.split("_", 1)[1], + created_at=run.created_at, + started_at=run.started_at, + completed_at=run.completed_at, + failed_at=run.failed_at, + cancelled_at=run.cancelled_at, + tried_cancelling_at=run.tried_cancelling_at, + expires_at=run.expires_at, + last_error=run.last_error.model_dump() + if run.last_error + else None, + ) + session.add(db_run) + + # Migrate events from JSONL + events_dir = data_dir / "events" + if events_dir.exists(): + for thread_dir in events_dir.iterdir(): + if thread_dir.is_dir(): + thread_id = thread_dir.name.split("_", 1)[1] # Remove prefix + for jsonl_file in thread_dir.glob("*.jsonl"): + run_id = jsonl_file.stem.split("_", 1)[1] # Remove prefix + sequence_num = 0 + with open(jsonl_file) as f: + for line in f: + event_json = line.strip() + if event_json: + # Parse and insert event + try: + event = Event.model_validate_json(event_json) + db_event = EventModel( + run_id=run_id, + thread_id=thread_id, + sequence_num=sequence_num, + event_type=event.event, + event_data=event_json, + created_at=now(), + ) + session.add(db_event) + sequence_num += 1 + except Exception: + # Skip invalid events + continue + + session.commit() + + # After successful migration, rename JSON directories + for dir_name in [ + "assistants", + "threads", + "messages", + "runs", + "files", + "mcp_configs", + "workflows", + "events", + ]: + dir_path = data_dir / dir_name + if dir_path.exists(): + dir_path.rename(data_dir / f"{dir_name}.migrated") + + except Exception: + session.rollback() + raise + finally: + session.close() + session.close() diff --git a/tests/integration/chat/api/conftest.py b/tests/integration/chat/api/conftest.py index a9272840..69dd5a0d 100644 --- a/tests/integration/chat/api/conftest.py +++ b/tests/integration/chat/api/conftest.py @@ -2,14 +2,18 @@ import tempfile import uuid +from contextlib import contextmanager from pathlib import Path +from typing import Generator import pytest -from fastapi import FastAPI -from fastapi.testclient import TestClient - from askui.chat.api.app import app +from askui.chat.api.db.base import Base from askui.chat.api.files.service import FileService +from fastapi import APIRouter, FastAPI +from fastapi.testclient import TestClient +from sqlalchemy import create_engine +from sqlalchemy.orm import Session, sessionmaker @pytest.fixture @@ -43,6 +47,329 @@ def test_headers(test_workspace_id: str) -> dict[str, str]: return {"askui-workspace": test_workspace_id} +@pytest.fixture +def test_db_engine(): + """Create in-memory SQLite database.""" + # Import all models to register them with Base.metadata + + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + yield engine + engine.dispose() + + +@pytest.fixture +def test_session_factory(test_db_engine): + """Create session factory for testing.""" + SessionLocal = sessionmaker(bind=test_db_engine, expire_on_commit=False) + + @contextmanager + def session_factory() -> Generator[Session, None, None]: + session = SessionLocal() + try: + yield session + finally: + session.close() + + return session_factory + + +@pytest.fixture +def test_app_with_db(test_db_engine): + """Create a test app with database tables created.""" + from askui.chat.api.db.base import Base + from askui.chat.api.dependencies import SetEnvFromHeadersDep + + # Import all models to register them with Base.metadata + from askui.chat.api.assistants.models import AssistantModel + from askui.chat.api.threads.models import ThreadModel + from askui.chat.api.messages.models import MessageModel + from askui.chat.api.runs.models import RunModel + from askui.chat.api.files.models import FileModel + from askui.chat.api.mcp_configs.models import McpConfigModel + from askui.chat.api.workflows.models import WorkflowModel, WorkflowTagModel + from askui.chat.api.runs.events.models import EventModel + from askui.chat.api.migrations.models import MigrationVersionModel + + # Create tables in the test database + Base.metadata.create_all(test_db_engine) + + # Create a new app instance without lifespan + test_app = FastAPI( + title="AskUI Chat API", + version="1.0.0", + dependencies=[SetEnvFromHeadersDep], + ) + + # Import and include all routers + from askui.chat.api.assistants.router import router as assistants_router + from askui.chat.api.files.router import router as files_router + from askui.chat.api.health.router import router as health_router + from askui.chat.api.mcp_configs.router import router as mcp_configs_router + from askui.chat.api.messages.router import router as messages_router + from askui.chat.api.runs.router import router as runs_router + from askui.chat.api.threads.router import router as threads_router + from askui.chat.api.workflows.router import router as workflows_router + + v1_router = APIRouter(prefix="/v1") + v1_router.include_router(assistants_router) + v1_router.include_router(threads_router) + v1_router.include_router(messages_router) + v1_router.include_router(runs_router) + v1_router.include_router(mcp_configs_router) + v1_router.include_router(files_router) + v1_router.include_router(workflows_router) + v1_router.include_router(health_router) + test_app.include_router(v1_router) + + return test_app + + +@pytest.fixture +def test_client_with_db(test_app_with_db, test_db_engine): + """Get test client with database.""" + # Import all models to register them with Base.metadata + from askui.chat.api.assistants.models import AssistantModel + from askui.chat.api.threads.models import ThreadModel + from askui.chat.api.messages.models import MessageModel + from askui.chat.api.runs.models import RunModel + from askui.chat.api.files.models import FileModel + from askui.chat.api.mcp_configs.models import McpConfigModel + from askui.chat.api.workflows.models import WorkflowModel, WorkflowTagModel + from askui.chat.api.runs.events.models import EventModel + from askui.chat.api.migrations.models import MigrationVersionModel + + # Ensure tables are created + Base.metadata.create_all(test_db_engine) + + SessionLocal = sessionmaker(bind=test_db_engine, expire_on_commit=False) + + @contextmanager + def session_factory() -> Generator[Session, None, None]: + session = SessionLocal() + try: + yield session + finally: + session.close() + + from askui.chat.api.assistants.dependencies import get_assistant_service + from askui.chat.api.assistants.service import AssistantService + from askui.chat.api.dependencies import get_session_factory_dep, get_settings + from askui.chat.api.files.dependencies import get_file_service + from askui.chat.api.files.service import FileService + from askui.chat.api.mcp_configs.dependencies import get_mcp_config_service + from askui.chat.api.mcp_configs.service import McpConfigService + from askui.chat.api.messages.dependencies import get_message_service + from askui.chat.api.messages.service import MessageService + from askui.chat.api.runs.dependencies import get_runs_service + from askui.chat.api.runs.service import RunService + from askui.chat.api.threads.dependencies import get_thread_service + from askui.chat.api.threads.service import ThreadService + from askui.chat.api.workflows.dependencies import get_workflow_service + from askui.chat.api.workflows.service import WorkflowService + + def get_test_assistant_service() -> AssistantService: + return AssistantService(session_factory) + + def get_test_file_service() -> FileService: + return FileService(session_factory) + + def get_test_mcp_config_service() -> McpConfigService: + return McpConfigService(session_factory, Path.cwd(), []) + + def get_test_message_service() -> MessageService: + return MessageService(session_factory) + + def get_test_runs_service() -> RunService: + from askui.chat.api.assistants.service import AssistantService + from askui.chat.api.mcp_clients.manager import McpClientManagerManager + from askui.chat.api.messages.chat_history_manager import ChatHistoryManager + from askui.chat.api.settings import Settings + + mock_assistant_service = AssistantService(session_factory) + mock_mcp_client_manager_manager = McpClientManagerManager() + mock_chat_history_manager = ChatHistoryManager() + mock_settings = Settings() + + return RunService( + session_factory=session_factory, + assistant_service=mock_assistant_service, + mcp_client_manager_manager=mock_mcp_client_manager_manager, + chat_history_manager=mock_chat_history_manager, + settings=mock_settings, + ) + + def get_test_thread_service() -> ThreadService: + from askui.chat.api.messages.service import MessageService + from askui.chat.api.runs.service import RunService + + mock_message_service = MessageService(session_factory) + mock_run_service = RunService( + session_factory=session_factory, + assistant_service=AssistantService(session_factory), + mcp_client_manager_manager=McpClientManagerManager(), + chat_history_manager=ChatHistoryManager(), + settings=Settings(), + ) + + return ThreadService( + session_factory=session_factory, + message_service=mock_message_service, + run_service=mock_run_service, + ) + + def get_test_workflow_service() -> WorkflowService: + return WorkflowService(session_factory) + + def get_test_session_factory(): + return session_factory + + def get_test_settings(): + from askui.chat.api.settings import Settings + settings = Settings() + # Override the database URL to use the test database + settings.db.url = f"sqlite:///:memory:" + return settings + + # Override all dependencies + test_app_with_db.dependency_overrides[get_settings] = get_test_settings + test_app_with_db.dependency_overrides[get_session_factory_dep] = ( + get_test_session_factory + ) + test_app_with_db.dependency_overrides[get_assistant_service] = ( + get_test_assistant_service + ) + test_app_with_db.dependency_overrides[get_file_service] = get_test_file_service + test_app_with_db.dependency_overrides[get_mcp_config_service] = ( + get_test_mcp_config_service + ) + test_app_with_db.dependency_overrides[get_message_service] = ( + get_test_message_service + ) + test_app_with_db.dependency_overrides[get_runs_service] = get_test_runs_service + test_app_with_db.dependency_overrides[get_thread_service] = get_test_thread_service + test_app_with_db.dependency_overrides[get_workflow_service] = ( + get_test_workflow_service + ) + + client = TestClient(test_app_with_db) + yield client + test_app_with_db.dependency_overrides.clear() + + +@pytest.fixture +def test_client_and_session_factory(test_app_with_db, test_db_engine): + """Get test client and session factory that use the same database.""" + # Ensure tables are created + Base.metadata.create_all(test_db_engine) + + SessionLocal = sessionmaker(bind=test_db_engine, expire_on_commit=False) + + @contextmanager + def session_factory() -> Generator[Session, None, None]: + session = SessionLocal() + try: + yield session + finally: + session.close() + + from askui.chat.api.assistants.dependencies import get_assistant_service + from askui.chat.api.assistants.service import AssistantService + from askui.chat.api.dependencies import get_session_factory_dep, get_settings + from askui.chat.api.files.dependencies import get_file_service + from askui.chat.api.files.service import FileService + from askui.chat.api.mcp_configs.dependencies import get_mcp_config_service + from askui.chat.api.mcp_configs.service import McpConfigService + from askui.chat.api.messages.dependencies import get_message_service + from askui.chat.api.messages.service import MessageService + from askui.chat.api.runs.dependencies import get_runs_service + from askui.chat.api.runs.service import RunService + from askui.chat.api.threads.dependencies import get_thread_service + from askui.chat.api.threads.service import ThreadService + from askui.chat.api.workflows.dependencies import get_workflow_service + from askui.chat.api.workflows.service import WorkflowService + + def get_test_assistant_service() -> AssistantService: + return AssistantService(session_factory) + + def get_test_file_service() -> FileService: + return FileService(session_factory) + + def get_test_mcp_config_service() -> McpConfigService: + return McpConfigService(session_factory, Path.cwd(), []) + + def get_test_message_service() -> MessageService: + return MessageService(session_factory) + + def get_test_runs_service() -> RunService: + from askui.chat.api.assistants.service import AssistantService + from askui.chat.api.mcp_clients.manager import McpClientManagerManager + from askui.chat.api.messages.chat_history_manager import ChatHistoryManager + from askui.chat.api.settings import Settings + + mock_assistant_service = AssistantService(session_factory) + mock_mcp_client_manager_manager = McpClientManagerManager() + mock_chat_history_manager = ChatHistoryManager() + mock_settings = Settings() + + return RunService( + session_factory=session_factory, + assistant_service=mock_assistant_service, + mcp_client_manager_manager=mock_mcp_client_manager_manager, + chat_history_manager=mock_chat_history_manager, + settings=mock_settings, + ) + + def get_test_thread_service() -> ThreadService: + from askui.chat.api.messages.service import MessageService + from askui.chat.api.runs.service import RunService + + mock_message_service = MessageService(session_factory) + mock_run_service = RunService( + session_factory=session_factory, + assistant_service=AssistantService(session_factory), + mcp_client_manager_manager=McpClientManagerManager(), + chat_history_manager=ChatHistoryManager(), + settings=Settings(), + ) + + return ThreadService( + session_factory=session_factory, + message_service=mock_message_service, + run_service=mock_run_service, + ) + + def get_test_workflow_service() -> WorkflowService: + return WorkflowService(session_factory) + + def get_test_session_factory(): + return session_factory + + # Override all dependencies + test_app_with_db.dependency_overrides[get_session_factory_dep] = ( + get_test_session_factory + ) + test_app_with_db.dependency_overrides[get_assistant_service] = ( + get_test_assistant_service + ) + test_app_with_db.dependency_overrides[get_file_service] = get_test_file_service + test_app_with_db.dependency_overrides[get_mcp_config_service] = ( + get_test_mcp_config_service + ) + test_app_with_db.dependency_overrides[get_message_service] = ( + get_test_message_service + ) + test_app_with_db.dependency_overrides[get_runs_service] = get_test_runs_service + test_app_with_db.dependency_overrides[get_thread_service] = get_test_thread_service + test_app_with_db.dependency_overrides[get_workflow_service] = ( + get_test_workflow_service + ) + + client = TestClient(test_app_with_db) + yield client, session_factory + test_app_with_db.dependency_overrides.clear() + + @pytest.fixture def mock_file_service(temp_workspace_dir: Path) -> FileService: """Create a mock file service with temporary workspace.""" diff --git a/tests/integration/chat/api/test_assistants.py b/tests/integration/chat/api/test_assistants.py index 6c60b64a..ad34f99a 100644 --- a/tests/integration/chat/api/test_assistants.py +++ b/tests/integration/chat/api/test_assistants.py @@ -1,315 +1,201 @@ """Integration tests for the assistants API endpoints.""" -import tempfile -from pathlib import Path +from datetime import datetime, timezone +from askui.chat.api.assistants.models import AssistantModel from fastapi import status from fastapi.testclient import TestClient -from askui.chat.api.assistants.models import Assistant -from askui.chat.api.assistants.service import AssistantService - class TestAssistantsAPI: """Test suite for the assistants API endpoints.""" - def test_list_assistants_empty(self, test_headers: dict[str, str]) -> None: + def test_list_assistants_empty( + self, test_client_with_db: TestClient, test_headers: dict[str, str] + ) -> None: """Test listing assistants when no assistants exist.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - - # Create a test app with overridden dependencies - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + response = test_client_with_db.get("/v1/assistants", headers=test_headers) - try: - with TestClient(app) as client: - response = client.get("/v1/assistants", headers=test_headers) - - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["object"] == "list" - assert data["data"] == [] - assert data["has_more"] is False - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["object"] == "list" + assert data["data"] == [] + assert data["has_more"] is False + assert data["first_id"] is None + assert data["last_id"] is None def test_list_assistants_with_assistants( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test listing assistants when assistants exist.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - # Create a mock assistant - workspace_id = test_headers["askui-workspace"] - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Test Assistant", - description="A test assistant", - avatar="test_avatar.png", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() - ) - - # Create a test app with overridden dependencies - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Test Assistant", + description="A test assistant", + avatar="test_avatar.png", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - response = client.get("/v1/assistants", headers=test_headers) + response = test_client.get("/v1/assistants", headers=test_headers) - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["object"] == "list" - assert len(data["data"]) == 1 - assert data["data"][0]["id"] == "asst_test123" - assert data["data"][0]["name"] == "Test Assistant" - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["object"] == "list" + assert len(data["data"]) == 1 + assert data["data"][0]["id"] == "asst_test123" + assert data["data"][0]["name"] == "Test Assistant" def test_list_assistants_with_pagination( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test listing assistants with pagination parameters.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - # Create multiple mock assistants - workspace_id = test_headers["askui-workspace"] - for i in range(5): - mock_assistant = Assistant( - id=f"asst_test{i}", - object="assistant", - created_at=1234567890 + i, - name=f"Test Assistant {i}", - description=f"Test assistant {i}", - workspace_id=workspace_id, - ) - (assistants_dir / f"asst_test{i}.json").write_text( - mock_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create multiple test assistants in the database + with session_factory() as session: + for i in range(5): + db_assistant = AssistantModel( + id=f"asst_test{i}", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name=f"Test Assistant {i}", + description=f"Test assistant {i}", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - response = client.get("/v1/assistants?limit=3", headers=test_headers) + response = test_client.get("/v1/assistants?limit=3", headers=test_headers) - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert len(data["data"]) == 3 - assert data["has_more"] is True - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert len(data["data"]) == 3 + assert data["has_more"] is True - def test_create_assistant(self, test_headers: dict[str, str]) -> None: + def test_create_assistant( + self, test_client_with_db: TestClient, test_headers: dict[str, str] + ) -> None: """Test creating a new assistant.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service - - try: - with TestClient(app) as client: - assistant_data = { - "name": "New Test Assistant", - "description": "A newly created test assistant", - "avatar": "new_avatar.png", - } - response = client.post( - "/v1/assistants", json=assistant_data, headers=test_headers - ) + assistant_data = { + "name": "New Test Assistant", + "description": "A newly created test assistant", + "avatar": "new_avatar.png", + } + response = test_client_with_db.post( + "/v1/assistants", json=assistant_data, headers=test_headers + ) - assert response.status_code == status.HTTP_201_CREATED - data = response.json() - assert data["name"] == "New Test Assistant" - assert data["description"] == "A newly created test assistant" - assert data["avatar"] == "new_avatar.png" - assert data["object"] == "assistant" - assert "id" in data - assert "created_at" in data - finally: - app.dependency_overrides.clear() - - def test_create_assistant_minimal(self, test_headers: dict[str, str]) -> None: + assert response.status_code == status.HTTP_201_CREATED + data = response.json() + assert data["name"] == "New Test Assistant" + assert data["description"] == "A newly created test assistant" + assert data["avatar"] == "new_avatar.png" + assert data["object"] == "assistant" + assert "id" in data + assert "created_at" in data + + def test_create_assistant_minimal( + self, test_client_with_db: TestClient, test_headers: dict[str, str] + ) -> None: """Test creating an assistant with minimal data.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service - - try: - with TestClient(app) as client: - response = client.post("/v1/assistants", json={}, headers=test_headers) + response = test_client_with_db.post( + "/v1/assistants", json={}, headers=test_headers + ) - assert response.status_code == status.HTTP_201_CREATED - data = response.json() - assert data["object"] == "assistant" - assert data["name"] is None - assert data["description"] is None - assert data["avatar"] is None - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_201_CREATED + data = response.json() + assert data["object"] == "assistant" + assert data["name"] is None + assert data["description"] is None + assert data["avatar"] is None def test_create_assistant_with_tools_and_system( - self, test_headers: dict[str, str] + self, test_client_with_db: TestClient, test_headers: dict[str, str] ) -> None: """Test creating a new assistant with tools and system prompt.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - - # Create a test app with overridden dependencies - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service - - try: - with TestClient(app) as client: - response = client.post( - "/v1/assistants", - headers=test_headers, - json={ - "name": "Custom Assistant", - "description": "A custom assistant with tools", - "tools": ["tool1", "tool2", "tool3"], - "system": "You are a helpful custom assistant.", - }, - ) + response = test_client_with_db.post( + "/v1/assistants", + headers=test_headers, + json={ + "name": "Custom Assistant", + "description": "A custom assistant with tools", + "tools": ["tool1", "tool2", "tool3"], + "system": "You are a helpful custom assistant.", + }, + ) - assert response.status_code == status.HTTP_201_CREATED - data = response.json() - assert data["name"] == "Custom Assistant" - assert data["description"] == "A custom assistant with tools" - assert data["tools"] == ["tool1", "tool2", "tool3"] - assert data["system"] == "You are a helpful custom assistant." - assert "id" in data - assert "created_at" in data - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_201_CREATED + data = response.json() + assert data["name"] == "Custom Assistant" + assert data["description"] == "A custom assistant with tools" + assert data["tools"] == ["tool1", "tool2", "tool3"] + assert data["system"] == "You are a helpful custom assistant." + assert "id" in data + assert "created_at" in data def test_create_assistant_with_empty_tools( - self, test_headers: dict[str, str] + self, test_client_with_db: TestClient, test_headers: dict[str, str] ) -> None: """Test creating a new assistant with empty tools list.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - - # Create a test app with overridden dependencies - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service - - try: - with TestClient(app) as client: - response = client.post( - "/v1/assistants", - headers=test_headers, - json={ - "name": "Empty Tools Assistant", - "tools": [], - }, - ) - - assert response.status_code == status.HTTP_201_CREATED - data = response.json() - assert data["name"] == "Empty Tools Assistant" - assert data["tools"] == [] - assert "id" in data - assert "created_at" in data - finally: - app.dependency_overrides.clear() - - def test_retrieve_assistant(self, test_headers: dict[str, str]) -> None: - """Test retrieving an existing assistant.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Test Assistant", - description="A test assistant", - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() + response = test_client_with_db.post( + "/v1/assistants", + headers=test_headers, + json={ + "name": "Empty Tools Assistant", + "tools": [], + }, ) - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) + assert response.status_code == status.HTTP_201_CREATED + data = response.json() + assert data["name"] == "Empty Tools Assistant" + assert data["tools"] == [] + assert "id" in data + assert "created_at" in data - app.dependency_overrides[get_assistant_service] = override_assistant_service + def test_retrieve_assistant( + self, test_client_and_session_factory, test_headers: dict[str, str] + ) -> None: + """Test retrieving an existing assistant.""" + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Test Assistant", + description="A test assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - response = client.get( - "/v1/assistants/asst_test123", headers=test_headers - ) + response = test_client.get("/v1/assistants/asst_test123", headers=test_headers) - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["id"] == "asst_test123" - assert data["name"] == "Test Assistant" - assert data["description"] == "A test assistant" - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["id"] == "asst_test123" + assert data["name"] == "Test Assistant" + assert data["description"] == "A test assistant" def test_retrieve_assistant_not_found( - self, test_client: TestClient, test_headers: dict[str, str] + self, test_client_with_db: TestClient, test_headers: dict[str, str] ) -> None: """Test retrieving a non-existent assistant.""" - response = test_client.get( + response = test_client_with_db.get( "/v1/assistants/asst_nonexistent123", headers=test_headers ) @@ -317,469 +203,347 @@ def test_retrieve_assistant_not_found( data = response.json() assert "detail" in data - def test_modify_assistant(self, test_headers: dict[str, str]) -> None: + def test_modify_assistant( + self, test_client_and_session_factory, test_headers: dict[str, str] + ) -> None: """Test modifying an existing assistant.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - workspace_id = test_headers["askui-workspace"] - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Original Name", - description="Original description", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Original Name", + description="Original description", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - modify_data = { - "name": "Modified Name", - "description": "Modified description", - } - response = client.post( - "/v1/assistants/asst_test123", - json=modify_data, - headers=test_headers, - ) + modify_data = { + "name": "Modified Name", + "description": "Modified description", + } + response = test_client.post( + "/v1/assistants/asst_test123", + json=modify_data, + headers=test_headers, + ) - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["name"] == "Modified Name" - assert data["description"] == "Modified description" - assert data["id"] == "asst_test123" - assert data["created_at"] == 1234567890 - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["name"] == "Modified Name" + assert data["description"] == "Modified description" + assert data["id"] == "asst_test123" def test_modify_assistant_with_tools_and_system( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test modifying an assistant with tools and system prompt.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - workspace_id = test_headers["askui-workspace"] - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Original Name", - description="Original description", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Original Name", + description="Original description", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() + + modify_data = { + "name": "Modified Name", + "tools": ["new_tool1", "new_tool2"], + "system": "You are a modified custom assistant.", + } + response = test_client.post( + "/v1/assistants/asst_test123", + json=modify_data, + headers=test_headers, ) - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service - - try: - with TestClient(app) as client: - modify_data = { - "name": "Modified Name", - "tools": ["new_tool1", "new_tool2"], - "system": "You are a modified custom assistant.", - } - response = client.post( - "/v1/assistants/asst_test123", - json=modify_data, - headers=test_headers, - ) + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["name"] == "Modified Name" + assert data["tools"] == ["new_tool1", "new_tool2"] + assert data["system"] == "You are a modified custom assistant." + assert data["id"] == "asst_test123" - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["name"] == "Modified Name" - assert data["tools"] == ["new_tool1", "new_tool2"] - assert data["system"] == "You are a modified custom assistant." - assert data["id"] == "asst_test123" - assert data["created_at"] == 1234567890 - finally: - app.dependency_overrides.clear() - - def test_modify_assistant_partial(self, test_headers: dict[str, str]) -> None: + def test_modify_assistant_partial( + self, test_client_and_session_factory, test_headers: dict[str, str] + ) -> None: """Test modifying an assistant with partial data.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - workspace_id = test_headers["askui-workspace"] - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Original Name", - description="Original description", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Original Name", + description="Original description", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - modify_data = {"name": "Only Name Modified"} - response = client.post( - "/v1/assistants/asst_test123", - json=modify_data, - headers=test_headers, - ) + modify_data = {"name": "Only Name Modified"} + response = test_client.post( + "/v1/assistants/asst_test123", + json=modify_data, + headers=test_headers, + ) - assert response.status_code == status.HTTP_200_OK - data = response.json() - assert data["name"] == "Only Name Modified" - assert data["description"] == "Original description" # Unchanged - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_200_OK + data = response.json() + assert data["name"] == "Only Name Modified" + assert data["description"] == "Original description" # Unchanged def test_modify_assistant_not_found( - self, test_client: TestClient, test_headers: dict[str, str] + self, test_client_with_db: TestClient, test_headers: dict[str, str] ) -> None: """Test modifying a non-existent assistant.""" modify_data = {"name": "Modified Name"} - response = test_client.post( + response = test_client_with_db.post( "/v1/assistants/asst_nonexistent123", json=modify_data, headers=test_headers ) assert response.status_code == status.HTTP_404_NOT_FOUND - def test_delete_assistant(self, test_headers: dict[str, str]) -> None: + def test_delete_assistant( + self, test_client_and_session_factory, test_headers: dict[str, str] + ) -> None: """Test deleting an existing assistant.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - workspace_id = test_headers["askui-workspace"] - mock_assistant = Assistant( - id="asst_test123", - object="assistant", - created_at=1234567890, - name="Test Assistant", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_test123.json").write_text( - mock_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a test assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_test123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Test Assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - response = client.delete( - "/v1/assistants/asst_test123", headers=test_headers - ) + response = test_client.delete( + "/v1/assistants/asst_test123", headers=test_headers + ) - assert response.status_code == status.HTTP_204_NO_CONTENT - assert response.content == b"" - finally: - app.dependency_overrides.clear() + assert response.status_code == status.HTTP_204_NO_CONTENT + assert response.content == b"" def test_delete_assistant_not_found( - self, test_client: TestClient, test_headers: dict[str, str] + self, test_client_with_db: TestClient, test_headers: dict[str, str] ) -> None: """Test deleting a non-existent assistant.""" - response = test_client.delete( + response = test_client_with_db.delete( "/v1/assistants/asst_nonexistent123", headers=test_headers ) assert response.status_code == status.HTTP_404_NOT_FOUND def test_modify_default_assistant_forbidden( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test that modifying a default assistant returns 403 Forbidden.""" - # Create a default assistant (no workspace_id) - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - default_assistant = Assistant( - id="asst_default123", - object="assistant", - created_at=1234567890, - name="Default Assistant", - description="This is a default assistant", - workspace_id=None, # No workspace_id = default - ) - (assistants_dir / "asst_default123.json").write_text( - default_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a default assistant (no workspace_id) in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_default123", + workspace_id=None, # No workspace_id = default + created_at=datetime.now(timezone.utc), + name="Default Assistant", + description="This is a default assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - # Try to modify the default assistant - response = client.post( - "/v1/assistants/asst_default123", - headers=test_headers, - json={"name": "Modified Name"}, - ) - assert response.status_code == 403 - assert "cannot be modified" in response.json()["detail"] - finally: - app.dependency_overrides.clear() + # Try to modify the default assistant + response = test_client.post( + "/v1/assistants/asst_default123", + headers=test_headers, + json={"name": "Modified Name"}, + ) + assert response.status_code == 403 + assert "cannot be modified" in response.json()["detail"] def test_delete_default_assistant_forbidden( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test that deleting a default assistant returns 403 Forbidden.""" - # Create a default assistant (no workspace_id) - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - default_assistant = Assistant( - id="asst_default456", - object="assistant", - created_at=1234567890, - name="Default Assistant", - description="This is a default assistant", - workspace_id=None, # No workspace_id = default - ) - (assistants_dir / "asst_default456.json").write_text( - default_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a default assistant (no workspace_id) in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_default456", + workspace_id=None, # No workspace_id = default + created_at=datetime.now(timezone.utc), + name="Default Assistant", + description="This is a default assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - # Try to delete the default assistant - response = client.delete( - "/v1/assistants/asst_default456", - headers=test_headers, - ) - assert response.status_code == 403 - assert "cannot be deleted" in response.json()["detail"] - finally: - app.dependency_overrides.clear() + # Try to delete the default assistant + response = test_client.delete( + "/v1/assistants/asst_default456", + headers=test_headers, + ) + assert response.status_code == 403 + assert "cannot be deleted" in response.json()["detail"] def test_list_assistants_includes_default_and_workspace( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test that listing assistants includes both default and workspace-scoped ones. """ - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - # Create a default assistant (no workspace_id) - default_assistant = Assistant( - id="asst_default789", - object="assistant", - created_at=1234567890, - name="Default Assistant", - description="This is a default assistant", - workspace_id=None, # No workspace_id = default - ) - (assistants_dir / "asst_default789.json").write_text( - default_assistant.model_dump_json() - ) - - # Create a workspace-scoped assistant - workspace_id = test_headers["askui-workspace"] - workspace_assistant = Assistant( - id="asst_workspace123", - object="assistant", - created_at=1234567890, - name="Workspace Assistant", - description="This is a workspace assistant", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_workspace123.json").write_text( - workspace_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a default assistant (no workspace_id) in the database + with session_factory() as session: + default_assistant = AssistantModel( + id="asst_default789", + workspace_id=None, # No workspace_id = default + created_at=datetime.now(timezone.utc), + name="Default Assistant", + description="This is a default assistant", + tools=[], + system=None, + ) + session.add(default_assistant) + + # Create a workspace-scoped assistant + workspace_assistant = AssistantModel( + id="asst_workspace123", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Workspace Assistant", + description="This is a workspace assistant", + tools=[], + system=None, + ) + session.add(workspace_assistant) + session.commit() - try: - with TestClient(app) as client: - # List assistants - should include both - response = client.get("/v1/assistants", headers=test_headers) - assert response.status_code == 200 + # List assistants - should include both + response = test_client.get("/v1/assistants", headers=test_headers) + assert response.status_code == 200 - data = response.json() - assistant_ids = [assistant["id"] for assistant in data["data"]] + data = response.json() + assistant_ids = [assistant["id"] for assistant in data["data"]] - # Should include both default and workspace assistants - assert "asst_default789" in assistant_ids - assert "asst_workspace123" in assistant_ids + # Should include both default and workspace assistants + assert "asst_default789" in assistant_ids + assert "asst_workspace123" in assistant_ids - # Verify workspace_id fields - default_assistant_data = next( - a for a in data["data"] if a["id"] == "asst_default789" - ) - workspace_assistant_data = next( - a for a in data["data"] if a["id"] == "asst_workspace123" - ) + # Verify workspace_id fields + default_assistant_data = next( + a for a in data["data"] if a["id"] == "asst_default789" + ) + workspace_assistant_data = next( + a for a in data["data"] if a["id"] == "asst_workspace123" + ) - assert default_assistant_data["workspace_id"] is None - assert workspace_assistant_data["workspace_id"] == workspace_id - finally: - app.dependency_overrides.clear() + assert default_assistant_data["workspace_id"] is None + assert ( + workspace_assistant_data["workspace_id"] == test_headers["askui-workspace"] + ) def test_retrieve_default_assistant_success( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test that retrieving a default assistant works.""" - # Create a default assistant (no workspace_id) - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - default_assistant = Assistant( - id="asst_defaultretrieve", - object="assistant", - created_at=1234567890, - name="Default Assistant", - description="This is a default assistant", - workspace_id=None, # No workspace_id = default - ) - (assistants_dir / "asst_defaultretrieve.json").write_text( - default_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a default assistant (no workspace_id) in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_defaultretrieve", + workspace_id=None, # No workspace_id = default + created_at=datetime.now(timezone.utc), + name="Default Assistant", + description="This is a default assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - # Retrieve the default assistant - response = client.get( - "/v1/assistants/asst_defaultretrieve", - headers=test_headers, - ) - assert response.status_code == 200 + # Retrieve the default assistant + response = test_client.get( + "/v1/assistants/asst_defaultretrieve", + headers=test_headers, + ) + assert response.status_code == 200 - data = response.json() - assert data["id"] == "asst_defaultretrieve" - assert data["workspace_id"] is None - finally: - app.dependency_overrides.clear() + data = response.json() + assert data["id"] == "asst_defaultretrieve" + assert data["workspace_id"] is None def test_workspace_scoped_assistant_operations_success( - self, test_headers: dict[str, str] + self, test_client_and_session_factory, test_headers: dict[str, str] ) -> None: """Test that workspace-scoped assistants can be modified and deleted.""" - temp_dir = tempfile.mkdtemp() - workspace_path = Path(temp_dir) - assistants_dir = workspace_path / "assistants" - assistants_dir.mkdir(parents=True, exist_ok=True) - - workspace_id = test_headers["askui-workspace"] - workspace_assistant = Assistant( - id="asst_workspaceops", - object="assistant", - created_at=1234567890, - name="Workspace Assistant", - description="This is a workspace assistant", - workspace_id=workspace_id, - ) - (assistants_dir / "asst_workspaceops.json").write_text( - workspace_assistant.model_dump_json() - ) - - from askui.chat.api.app import app - from askui.chat.api.assistants.dependencies import get_assistant_service - - def override_assistant_service() -> AssistantService: - return AssistantService(workspace_path) - - app.dependency_overrides[get_assistant_service] = override_assistant_service + test_client, session_factory = test_client_and_session_factory + + # Create a workspace-scoped assistant in the database + with session_factory() as session: + db_assistant = AssistantModel( + id="asst_workspaceops", + workspace_id=test_headers["askui-workspace"], + created_at=datetime.now(timezone.utc), + name="Workspace Assistant", + description="This is a workspace assistant", + tools=[], + system=None, + ) + session.add(db_assistant) + session.commit() - try: - with TestClient(app) as client: - # Modify the workspace assistant - response = client.post( - "/v1/assistants/asst_workspaceops", - headers=test_headers, - json={"name": "Modified Workspace Assistant"}, - ) - assert response.status_code == 200 + # Modify the workspace assistant + response = test_client.post( + "/v1/assistants/asst_workspaceops", + headers=test_headers, + json={"name": "Modified Workspace Assistant"}, + ) + assert response.status_code == 200 - data = response.json() - assert data["name"] == "Modified Workspace Assistant" - assert data["workspace_id"] == workspace_id + data = response.json() + assert data["name"] == "Modified Workspace Assistant" + assert data["workspace_id"] == test_headers["askui-workspace"] - # Delete the workspace assistant - response = client.delete( - "/v1/assistants/asst_workspaceops", - headers=test_headers, - ) - assert response.status_code == 204 + # Delete the workspace assistant + response = test_client.delete( + "/v1/assistants/asst_workspaceops", + headers=test_headers, + ) + assert response.status_code == 204 - # Verify it's deleted - response = client.get( - "/v1/assistants/asst_workspaceops", - headers=test_headers, - ) - assert response.status_code == 404 - finally: - app.dependency_overrides.clear() + # Verify it's deleted + response = test_client.get( + "/v1/assistants/asst_workspaceops", + headers=test_headers, + ) + assert response.status_code == 404 diff --git a/tests/integration/chat/api/test_assistants_service.py b/tests/integration/chat/api/test_assistants_service.py new file mode 100644 index 00000000..6211cff3 --- /dev/null +++ b/tests/integration/chat/api/test_assistants_service.py @@ -0,0 +1,268 @@ +"""Unit tests for the assistants service.""" + +import uuid +from contextlib import contextmanager + +import pytest +from askui.chat.api.assistants.schemas import ( + AssistantCreateParams, + AssistantModifyParams, +) +from askui.chat.api.assistants.service import AssistantService +from askui.chat.api.db.base import Base +from askui.utils.api_utils import ListQuery +from askui.utils.not_given import NOT_GIVEN +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker + + +@pytest.fixture +def test_db_engine(): + """Create in-memory SQLite database.""" + engine = create_engine("sqlite:///:memory:") + Base.metadata.create_all(engine) + yield engine + engine.dispose() + + +@pytest.fixture +def test_session_factory(test_db_engine): + """Create session factory for testing.""" + SessionLocal = sessionmaker(bind=test_db_engine, expire_on_commit=False) + + @contextmanager + def session_factory(): + session = SessionLocal() + try: + yield session + finally: + session.close() + + return session_factory + + +@pytest.fixture +def assistant_service(test_session_factory): + """Create assistant service for testing.""" + return AssistantService(test_session_factory) + + +@pytest.fixture +def test_workspace_id(): + """Get a test workspace ID.""" + return str(uuid.uuid4()) + + +class TestAssistantService: + """Test suite for the AssistantService.""" + + def test_create_assistant( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test creating a new assistant.""" + params = AssistantCreateParams( + name="Test Assistant", + description="A test assistant", + tools=["tool1", "tool2"], + system="You are a helpful assistant.", + ) + + assistant = assistant_service.create( + workspace_id=test_workspace_id, params=params + ) + + assert assistant.name == "Test Assistant" + assert assistant.description == "A test assistant" + assert assistant.tools == ["tool1", "tool2"] + assert assistant.system == "You are a helpful assistant." + assert str(assistant.workspace_id) == test_workspace_id + assert assistant.id.startswith("asst_") + assert assistant.created_at is not None + + def test_list_assistants_empty( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test listing assistants when no assistants exist.""" + query = ListQuery() + result = assistant_service.list_(workspace_id=test_workspace_id, query=query) + + assert len(result.data) == 0 + assert result.has_more is False + assert result.first_id is None + assert result.last_id is None + + def test_list_assistants_with_data( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test listing assistants when assistants exist.""" + # Create multiple assistants + for i in range(3): + params = AssistantCreateParams( + name=f"Assistant {i}", + description=f"Test assistant {i}", + tools=[], + system=None, + ) + assistant_service.create(workspace_id=test_workspace_id, params=params) + + query = ListQuery() + result = assistant_service.list_(workspace_id=test_workspace_id, query=query) + + assert len(result.data) == 3 + assert result.has_more is False + assert result.first_id is not None + assert result.last_id is not None + assert ( + result.data[0].name == "Assistant 2" + ) # Should be ordered by created_at desc (newest first) + + def test_retrieve_assistant( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test retrieving an existing assistant.""" + params = AssistantCreateParams( + name="Test Assistant", + description="A test assistant", + tools=[], + system=None, + ) + created_assistant = assistant_service.create( + workspace_id=test_workspace_id, params=params + ) + + retrieved_assistant = assistant_service.retrieve( + workspace_id=test_workspace_id, assistant_id=created_assistant.id + ) + + assert retrieved_assistant.id == created_assistant.id + assert retrieved_assistant.name == "Test Assistant" + assert retrieved_assistant.description == "A test assistant" + + def test_modify_assistant( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test modifying an existing assistant.""" + params = AssistantCreateParams( + name="Original Name", + description="Original description", + tools=[], + system=None, + ) + created_assistant = assistant_service.create( + workspace_id=test_workspace_id, params=params + ) + + modify_params = AssistantModifyParams( + name="Modified Name", + description="Modified description", + tools=["new_tool"], + system="You are a modified assistant.", + ) + + modified_assistant = assistant_service.modify( + workspace_id=test_workspace_id, + assistant_id=created_assistant.id, + params=modify_params, + ) + + assert modified_assistant.id == created_assistant.id + assert modified_assistant.name == "Modified Name" + assert modified_assistant.description == "Modified description" + assert modified_assistant.tools == ["new_tool"] + assert modified_assistant.system == "You are a modified assistant." + + def test_modify_assistant_partial( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test modifying an assistant with partial data.""" + params = AssistantCreateParams( + name="Original Name", + description="Original description", + tools=[], + system=None, + ) + created_assistant = assistant_service.create( + workspace_id=test_workspace_id, params=params + ) + + modify_params = AssistantModifyParams( + name="Modified Name", + description=NOT_GIVEN, + tools=NOT_GIVEN, + system=NOT_GIVEN, + ) + + modified_assistant = assistant_service.modify( + workspace_id=test_workspace_id, + assistant_id=created_assistant.id, + params=modify_params, + ) + + assert modified_assistant.id == created_assistant.id + assert modified_assistant.name == "Modified Name" + assert modified_assistant.description == "Original description" # Unchanged + assert modified_assistant.tools == [] # Unchanged + assert modified_assistant.system is None # Unchanged + + def test_delete_assistant( + self, assistant_service: AssistantService, test_workspace_id: str + ) -> None: + """Test deleting an existing assistant.""" + params = AssistantCreateParams( + name="Test Assistant", + description="A test assistant", + tools=[], + system=None, + ) + created_assistant = assistant_service.create( + workspace_id=test_workspace_id, params=params + ) + + # Delete the assistant + assistant_service.delete( + workspace_id=test_workspace_id, assistant_id=created_assistant.id + ) + + # Try to retrieve it - should raise NotFoundError + with pytest.raises(Exception): # Should be NotFoundError + assistant_service.retrieve( + workspace_id=test_workspace_id, assistant_id=created_assistant.id + ) + + def test_workspace_isolation(self, assistant_service: AssistantService) -> None: + """Test that assistants are isolated by workspace.""" + workspace1 = str(uuid.uuid4()) + workspace2 = str(uuid.uuid4()) + + # Create assistant in workspace1 + params1 = AssistantCreateParams( + name="Workspace 1 Assistant", tools=[], system=None + ) + assistant1 = assistant_service.create(workspace_id=workspace1, params=params1) + + # Create assistant in workspace2 + params2 = AssistantCreateParams( + name="Workspace 2 Assistant", tools=[], system=None + ) + assistant2 = assistant_service.create(workspace_id=workspace2, params=params2) + + # List assistants in workspace1 - should only see assistant1 + query = ListQuery() + result1 = assistant_service.list_(workspace_id=workspace1, query=query) + assert len(result1.data) == 1 + assert result1.data[0].id == assistant1.id + + # List assistants in workspace2 - should only see assistant2 + result2 = assistant_service.list_(workspace_id=workspace2, query=query) + assert len(result2.data) == 1 + assert result2.data[0].id == assistant2.id + + # Try to retrieve assistant1 from workspace2 - should fail + with pytest.raises(Exception): # Should be NotFoundError + assistant_service.retrieve( + workspace_id=workspace2, assistant_id=assistant1.id + ) + with pytest.raises(Exception): # Should be NotFoundError + assistant_service.retrieve( + workspace_id=workspace2, assistant_id=assistant1.id + ) diff --git a/tests/integration/chat/api/test_events_streaming.py b/tests/integration/chat/api/test_events_streaming.py new file mode 100644 index 00000000..49b71d71 --- /dev/null +++ b/tests/integration/chat/api/test_events_streaming.py @@ -0,0 +1,178 @@ +"""Integration tests for event streaming.""" + +import json + +import pytest +from fastapi.testclient import TestClient + + +def test_event_streaming_with_multiple_readers( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test event streaming with multiple readers.""" + # Create a thread first + thread_data = {"name": "Test Thread"} + response = test_client_with_db.post( + "/threads", json=thread_data, headers=test_headers + ) + assert response.status_code == 201 + thread_id = response.json()["id"] + + # Create an assistant + assistant_data = { + "name": "Test Assistant", + "description": "Test assistant", + "tools": [], + "system": "You are a test assistant", + } + response = test_client_with_db.post( + "/assistants", json=assistant_data, headers=test_headers + ) + assert response.status_code == 201 + assistant_id = response.json()["id"] + + # Create a run + run_data = {"assistant_id": assistant_id, "instructions": "Test instructions"} + response = test_client_with_db.post( + f"/threads/{thread_id}/runs", json=run_data, headers=test_headers + ) + assert response.status_code == 201 + run_id = response.json()["id"] + + # Start streaming events + response = test_client_with_db.get( + f"/threads/{thread_id}/runs/{run_id}/events/stream", headers=test_headers + ) + assert response.status_code == 200 + + # The response should be a streaming response + assert response.headers["content-type"] == "text/event-stream" + + # Read the stream content + stream_content = response.content.decode("utf-8") + + # Should contain event data + assert "data:" in stream_content + assert "event:" in stream_content + + +def test_event_streaming_with_cancellation( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test event streaming with run cancellation.""" + # Create a thread first + thread_data = {"name": "Test Thread"} + response = test_client_with_db.post( + "/threads", json=thread_data, headers=test_headers + ) + assert response.status_code == 201 + thread_id = response.json()["id"] + + # Create an assistant + assistant_data = { + "name": "Test Assistant", + "description": "Test assistant", + "tools": [], + "system": "You are a test assistant", + } + response = test_client_with_db.post( + "/assistants", json=assistant_data, headers=test_headers + ) + assert response.status_code == 201 + assistant_id = response.json()["id"] + + # Create a run + run_data = {"assistant_id": assistant_id, "instructions": "Test instructions"} + response = test_client_with_db.post( + f"/threads/{thread_id}/runs", json=run_data, headers=test_headers + ) + assert response.status_code == 201 + run_id = response.json()["id"] + + # Cancel the run + response = test_client_with_db.post( + f"/threads/{thread_id}/runs/{run_id}/cancel", headers=test_headers + ) + assert response.status_code == 200 + + # Start streaming events + response = test_client_with_db.get( + f"/threads/{thread_id}/runs/{run_id}/events/stream", headers=test_headers + ) + assert response.status_code == 200 + + # Read the stream content + stream_content = response.content.decode("utf-8") + + # Should contain cancellation event + assert "cancelled" in stream_content or "cancel" in stream_content.lower() + + +def test_event_streaming_with_error_handling( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test event streaming with error handling.""" + # Try to stream events for non-existent run + fake_thread_id = "thread_507f1f77bcf86cd799439011" + fake_run_id = "run_507f1f77bcf86cd799439012" + + response = test_client_with_db.get( + f"/threads/{fake_thread_id}/runs/{fake_run_id}/events/stream", + headers=test_headers, + ) + assert response.status_code == 404 + + +def test_event_streaming_content_format( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test that event streaming returns properly formatted content.""" + # Create a thread first + thread_data = {"name": "Test Thread"} + response = test_client_with_db.post( + "/threads", json=thread_data, headers=test_headers + ) + assert response.status_code == 201 + thread_id = response.json()["id"] + + # Create an assistant + assistant_data = { + "name": "Test Assistant", + "description": "Test assistant", + "tools": [], + "system": "You are a test assistant", + } + response = test_client_with_db.post( + "/assistants", json=assistant_data, headers=test_headers + ) + assert response.status_code == 201 + assistant_id = response.json()["id"] + + # Create a run + run_data = {"assistant_id": assistant_id, "instructions": "Test instructions"} + response = test_client_with_db.post( + f"/threads/{thread_id}/runs", json=run_data, headers=test_headers + ) + assert response.status_code == 201 + run_id = response.json()["id"] + + # Start streaming events + response = test_client_with_db.get( + f"/threads/{thread_id}/runs/{run_id}/events/stream", headers=test_headers + ) + assert response.status_code == 200 + + # Read the stream content + stream_content = response.content.decode("utf-8") + + # Should contain properly formatted SSE + lines = stream_content.strip().split("\n") + for line in lines: + if line.startswith("data:"): + # Should be valid JSON + data_content = line[5:].strip() + if data_content: + try: + json.loads(data_content) + except json.JSONDecodeError: + pytest.fail(f"Invalid JSON in event stream: {data_content}") diff --git a/tests/integration/chat/api/test_files.py b/tests/integration/chat/api/test_files.py index 4496794c..cb89ed2f 100644 --- a/tests/integration/chat/api/test_files.py +++ b/tests/integration/chat/api/test_files.py @@ -4,12 +4,11 @@ import tempfile from pathlib import Path +from askui.chat.api.files.schemas import File +from askui.chat.api.files.service import FileService from fastapi import status from fastapi.testclient import TestClient -from askui.chat.api.files.models import File -from askui.chat.api.files.service import FileService - class TestFilesAPI: """Test suite for the files API endpoints.""" diff --git a/tests/integration/chat/api/test_files_service.py b/tests/integration/chat/api/test_files_service.py index 49d221b1..f4565a7b 100644 --- a/tests/integration/chat/api/test_files_service.py +++ b/tests/integration/chat/api/test_files_service.py @@ -5,12 +5,11 @@ from unittest.mock import AsyncMock import pytest -from fastapi import UploadFile - -from askui.chat.api.files.models import File, FileCreateParams +from askui.chat.api.files.schemas import File, FileCreateParams from askui.chat.api.files.service import FileService from askui.chat.api.models import FileId from askui.utils.api_utils import ConflictError, FileTooLargeError, NotFoundError +from fastapi import UploadFile class TestFileService: diff --git a/tests/integration/chat/api/test_request_document_translator.py b/tests/integration/chat/api/test_request_document_translator.py index ba5f1597..c26654b5 100644 --- a/tests/integration/chat/api/test_request_document_translator.py +++ b/tests/integration/chat/api/test_request_document_translator.py @@ -6,14 +6,13 @@ from typing import Generator import pytest -from PIL import Image - from askui.chat.api.files.service import FileService -from askui.chat.api.messages.models import RequestDocumentBlockParam +from askui.chat.api.messages.schemas import RequestDocumentBlockParam from askui.chat.api.messages.translator import RequestDocumentBlockParamTranslator from askui.models.shared.agent_message_param import CacheControlEphemeralParam from askui.utils.excel_utils import OfficeDocumentSource from askui.utils.image_utils import ImageSource +from PIL import Image class TestRequestDocumentBlockParamTranslator: diff --git a/tests/integration/chat/api/test_runs.py b/tests/integration/chat/api/test_runs.py index ed06f357..33e92455 100644 --- a/tests/integration/chat/api/test_runs.py +++ b/tests/integration/chat/api/test_runs.py @@ -4,14 +4,13 @@ from pathlib import Path from unittest.mock import Mock -from fastapi import status -from fastapi.testclient import TestClient - from askui.chat.api.assistants.service import AssistantService from askui.chat.api.runs.models import Run from askui.chat.api.runs.service import RunService from askui.chat.api.threads.models import Thread from askui.chat.api.threads.service import ThreadService +from fastapi import status +from fastapi.testclient import TestClient def create_mock_mcp_client_manager_manager() -> Mock: @@ -45,12 +44,18 @@ def test_list_runs_empty(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -115,12 +120,18 @@ def test_list_runs_with_runs(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -186,12 +197,18 @@ def test_list_runs_with_pagination(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -308,12 +325,18 @@ def test_create_run_minimal(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -366,12 +389,18 @@ def test_create_run_streaming(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -413,12 +442,18 @@ def test_create_thread_and_run(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -472,12 +507,18 @@ def test_create_thread_and_run_minimal(self, test_headers: dict[str, str]) -> No from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -522,12 +563,18 @@ def test_create_thread_and_run_streaming( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -575,12 +622,18 @@ def test_create_thread_and_run_with_messages( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -641,12 +694,18 @@ def test_create_thread_and_run_validation_error( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -689,12 +748,18 @@ def test_create_thread_and_run_empty_thread( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -774,12 +839,18 @@ def test_retrieve_run(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -858,12 +929,18 @@ def test_cancel_run(self, test_headers: dict[str, str]) -> None: from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_assistant_service = Mock() @@ -946,20 +1023,26 @@ def test_create_run_with_custom_assistant( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_mcp_client_manager_manager = create_mock_mcp_client_manager_manager() from askui.chat.api.assistants.service import AssistantService return RunService( - base_dir=workspace_path, - assistant_service=AssistantService(workspace_path), + session_factory=mock_session_factory, + assistant_service=AssistantService(mock_session_factory), mcp_client_manager_manager=mock_mcp_client_manager_manager, chat_history_manager=Mock(), settings=Mock(), @@ -968,7 +1051,7 @@ def override_runs_service() -> RunService: def override_assistant_service() -> AssistantService: from askui.chat.api.assistants.service import AssistantService - return AssistantService(workspace_path) + return AssistantService(mock_session_factory) app.dependency_overrides[get_thread_service] = override_thread_service app.dependency_overrides[get_runs_service] = override_runs_service @@ -1032,20 +1115,26 @@ def test_create_run_with_custom_assistant_empty_tools( from askui.chat.api.runs.dependencies import get_runs_service from askui.chat.api.threads.dependencies import get_thread_service + mock_session_factory = Mock() + mock_session_factory.return_value.__enter__ = Mock(return_value=Mock()) + mock_session_factory.return_value.__exit__ = Mock(return_value=None) + def override_thread_service() -> ThreadService: from askui.chat.api.threads.service import ThreadService mock_message_service = Mock() mock_run_service = Mock() - return ThreadService(workspace_path, mock_message_service, mock_run_service) + return ThreadService( + mock_session_factory, mock_message_service, mock_run_service + ) def override_runs_service() -> RunService: mock_mcp_client_manager_manager = create_mock_mcp_client_manager_manager() from askui.chat.api.assistants.service import AssistantService return RunService( - base_dir=workspace_path, - assistant_service=AssistantService(workspace_path), + session_factory=mock_session_factory, + assistant_service=AssistantService(mock_session_factory), mcp_client_manager_manager=mock_mcp_client_manager_manager, chat_history_manager=Mock(), settings=Mock(), @@ -1054,7 +1143,7 @@ def override_runs_service() -> RunService: def override_assistant_service() -> AssistantService: from askui.chat.api.assistants.service import AssistantService - return AssistantService(workspace_path) + return AssistantService(mock_session_factory) app.dependency_overrides[get_thread_service] = override_thread_service app.dependency_overrides[get_runs_service] = override_runs_service diff --git a/tests/integration/chat/api/test_workflows_tags.py b/tests/integration/chat/api/test_workflows_tags.py new file mode 100644 index 00000000..fc0d5721 --- /dev/null +++ b/tests/integration/chat/api/test_workflows_tags.py @@ -0,0 +1,183 @@ +"""Integration tests for workflow tag filtering.""" + +from fastapi.testclient import TestClient + + +def test_list_workflows_filter_by_single_tag( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test filtering workflows by a single tag.""" + # Create workflows with different tags + workflow1_data = { + "name": "Workflow 1", + "description": "First workflow", + "tags": ["automation", "testing"], + } + workflow2_data = { + "name": "Workflow 2", + "description": "Second workflow", + "tags": ["testing", "deployment"], + } + workflow3_data = { + "name": "Workflow 3", + "description": "Third workflow", + "tags": ["automation", "deployment"], + } + + # Create workflows + response1 = test_client_with_db.post( + "/workflows", json=workflow1_data, headers=test_headers + ) + assert response1.status_code == 201 + + response2 = test_client_with_db.post( + "/workflows", json=workflow2_data, headers=test_headers + ) + assert response2.status_code == 201 + + response3 = test_client_with_db.post( + "/workflows", json=workflow3_data, headers=test_headers + ) + assert response3.status_code == 201 + + # Filter by "automation" tag + response = test_client_with_db.get( + "/workflows?tags=automation", headers=test_headers + ) + assert response.status_code == 200 + + data = response.json() + assert len(data["data"]) == 2 + workflow_names = [w["name"] for w in data["data"]] + assert "Workflow 1" in workflow_names + assert "Workflow 3" in workflow_names + assert "Workflow 2" not in workflow_names + + +def test_list_workflows_filter_by_multiple_tags( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test filtering workflows by multiple tags (OR logic).""" + # Create workflows with different tags + workflow1_data = { + "name": "Workflow 1", + "description": "First workflow", + "tags": ["automation", "testing"], + } + workflow2_data = { + "name": "Workflow 2", + "description": "Second workflow", + "tags": ["testing", "deployment"], + } + workflow3_data = { + "name": "Workflow 3", + "description": "Third workflow", + "tags": ["automation", "deployment"], + } + + # Create workflows + test_client_with_db.post("/workflows", json=workflow1_data, headers=test_headers) + test_client_with_db.post("/workflows", json=workflow2_data, headers=test_headers) + test_client_with_db.post("/workflows", json=workflow3_data, headers=test_headers) + + # Filter by "automation" OR "deployment" tags + response = test_client_with_db.get( + "/workflows?tags=automation&tags=deployment", headers=test_headers + ) + assert response.status_code == 200 + + data = response.json() + assert len(data["data"]) == 3 # All workflows should match + workflow_names = [w["name"] for w in data["data"]] + assert "Workflow 1" in workflow_names + assert "Workflow 2" in workflow_names + assert "Workflow 3" in workflow_names + + +def test_list_workflows_tag_filtering_with_pagination( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test combining tag filtering with pagination.""" + # Create multiple workflows with automation tag + for i in range(5): + workflow_data = { + "name": f"Automation Workflow {i}", + "description": f"Workflow {i}", + "tags": ["automation"], + } + test_client_with_db.post("/workflows", json=workflow_data, headers=test_headers) + + # Create workflows with different tags + for i in range(3): + workflow_data = { + "name": f"Testing Workflow {i}", + "description": f"Workflow {i}", + "tags": ["testing"], + } + test_client_with_db.post("/workflows", json=workflow_data, headers=test_headers) + + # Filter by automation tag with limit + response = test_client_with_db.get( + "/workflows?tags=automation&limit=3", headers=test_headers + ) + assert response.status_code == 200 + + data = response.json() + assert len(data["data"]) == 3 + assert data["has_more"] is True + + # Verify all returned workflows have automation tag + for workflow in data["data"]: + assert "automation" in workflow["tags"] + + +def test_modify_workflow_update_tags( + test_client_with_db: TestClient, test_headers: dict[str, str] +): + """Test modifying workflow to change tags.""" + # Create workflow with initial tags + workflow_data = { + "name": "Test Workflow", + "description": "Test workflow", + "tags": ["automation", "testing"], + } + + response = test_client_with_db.post( + "/workflows", json=workflow_data, headers=test_headers + ) + assert response.status_code == 201 + workflow_id = response.json()["id"] + + # Modify workflow to change tags + modify_data = { + "name": "Updated Workflow", + "description": "Updated description", + "tags": ["deployment", "production"], + } + + response = test_client_with_db.post( + f"/workflows/{workflow_id}", json=modify_data, headers=test_headers + ) + assert response.status_code == 200 + + updated_workflow = response.json() + assert updated_workflow["name"] == "Updated Workflow" + assert updated_workflow["description"] == "Updated description" + assert set(updated_workflow["tags"]) == {"deployment", "production"} + + # Verify old tags are no longer associated + response = test_client_with_db.get( + "/workflows?tags=automation", headers=test_headers + ) + assert response.status_code == 200 + data = response.json() + assert len(data["data"]) == 0 + + # Verify new tags are associated + response = test_client_with_db.get( + "/workflows?tags=deployment", headers=test_headers + ) + assert response.status_code == 200 + data = response.json() + assert len(data["data"]) == 1 + assert data["data"][0]["id"] == workflow_id diff --git a/tests/integration/chat/migrations/__init__.py b/tests/integration/chat/migrations/__init__.py new file mode 100644 index 00000000..cfe2a82c --- /dev/null +++ b/tests/integration/chat/migrations/__init__.py @@ -0,0 +1 @@ +"""Migration tests package.""" diff --git a/tests/integration/chat/migrations/test_migration_runner.py b/tests/integration/chat/migrations/test_migration_runner.py new file mode 100644 index 00000000..77cb729f --- /dev/null +++ b/tests/integration/chat/migrations/test_migration_runner.py @@ -0,0 +1,189 @@ +"""Integration tests for database migrations.""" + +from datetime import datetime, timezone +from pathlib import Path + +from askui.chat.api.assistants.models import Assistant +from askui.chat.api.workflows.models import Workflow +from askui.chat.migrations import MigrationRunner + + +def test_migration_from_empty_db(tmp_path: Path): + """Test migration from empty database.""" + db_path = tmp_path / "test.db" + data_dir = tmp_path / "data" + data_dir.mkdir() + + runner = MigrationRunner(db_path) + + # Should need migration for empty database (version 0 -> 2) + assert runner.should_migrate(data_dir) + + # Run migration + runner.migrate(data_dir) + + # Check that schema version was recorded + assert runner.get_current_version() == 2 + + +def test_migration_with_json_files(tmp_path: Path): + """Test migration with existing JSON files.""" + db_path = tmp_path / "test.db" + data_dir = tmp_path / "data" + data_dir.mkdir() + + # Create sample JSON files + assistants_dir = data_dir / "assistants" + assistants_dir.mkdir() + + assistant = Assistant( + id="asst_507f1f77bcf86cd799439011", + workspace_id="test-workspace", + created_at=datetime.now(timezone.utc), + name="Test Assistant", + description="Test description", + avatar=None, + tools=["test_tool"], + system="Test system", + ) + + assistant_file = assistants_dir / "asst_507f1f77bcf86cd799439011.json" + assistant_file.write_text(assistant.model_dump_json()) + + # Create workflow with tags + workflows_dir = data_dir / "workflows" + workflows_dir.mkdir() + + workflow = Workflow( + id="workflow_507f1f77bcf86cd799439012", + workspace_id="test-workspace", + created_at=datetime.now(timezone.utc), + name="Test Workflow", + description="Test description", + tags=["automation", "testing"], + ) + + workflow_file = workflows_dir / "workflow_507f1f77bcf86cd799439012.json" + workflow_file.write_text(workflow.model_dump_json()) + + runner = MigrationRunner(db_path) + + # Should need migration + assert runner.should_migrate(data_dir) + + # Run migration + runner.migrate(data_dir) + + # Check that schema version was recorded + assert runner.get_current_version() == 2 + + # Check that JSON directories were renamed + assert (data_dir / "assistants.migrated").exists() + assert (data_dir / "workflows.migrated").exists() + assert not (data_dir / "assistants").exists() + assert not (data_dir / "workflows").exists() + + +def test_migration_idempotent(tmp_path: Path): + """Test that migration is idempotent.""" + db_path = tmp_path / "test.db" + data_dir = tmp_path / "data" + data_dir.mkdir() + + runner = MigrationRunner(db_path) + + # Run migration twice + runner.migrate(data_dir) + version_after_first = runner.get_current_version() + + runner.migrate(data_dir) + version_after_second = runner.get_current_version() + + # Should be the same version + assert version_after_first == version_after_second == 2 + + +def test_workflow_tags_migration(tmp_path: Path): + """Test migration of workflow tags to separate table.""" + db_path = tmp_path / "test.db" + data_dir = tmp_path / "data" + data_dir.mkdir() + + # Create workflow with tags + workflows_dir = data_dir / "workflows" + workflows_dir.mkdir() + + workflow = Workflow( + id="workflow_507f1f77bcf86cd799439012", + workspace_id="test-workspace", + created_at=datetime.now(timezone.utc), + name="Test Workflow", + description="Test description", + tags=["automation", "testing", "deployment"], + ) + + workflow_file = workflows_dir / "workflow_507f1f77bcf86cd799439012.json" + workflow_file.write_text(workflow.model_dump_json()) + + runner = MigrationRunner(db_path) + runner.migrate(data_dir) + + # Check that tags were migrated to separate table + from sqlalchemy import create_engine, text + + engine = create_engine(f"sqlite:///{db_path}") + + with engine.connect() as conn: + # Check workflow_tags table + result = conn.execute( + text( + "SELECT * FROM workflow_tags WHERE workflow_id = '507f1f77bcf86cd799439012'" + ) + ) + tags = [row[2] for row in result] # tag column is index 2 + + assert len(tags) == 3 + assert "automation" in tags + assert "testing" in tags + assert "deployment" in tags + + +def test_objectid_prefix_preserved(tmp_path: Path): + """Test that ObjectId prefixes are preserved during migration.""" + db_path = tmp_path / "test.db" + data_dir = tmp_path / "data" + data_dir.mkdir() + + # Create assistant with prefixed ID + assistants_dir = data_dir / "assistants" + assistants_dir.mkdir() + + assistant = Assistant( + id="asst_507f1f77bcf86cd799439011", + workspace_id="test-workspace", + created_at=datetime.now(timezone.utc), + name="Test Assistant", + description="Test description", + avatar=None, + tools=["test_tool"], + system="Test system", + ) + + assistant_file = assistants_dir / "asst_507f1f77bcf86cd799439011.json" + assistant_file.write_text(assistant.model_dump_json()) + + runner = MigrationRunner(db_path) + runner.migrate(data_dir) + + # Check that ObjectId was stored without prefix in DB + from sqlalchemy import create_engine, text + + engine = create_engine(f"sqlite:///{db_path}") + + with engine.connect() as conn: + result = conn.execute( + text("SELECT id FROM assistants WHERE id = '507f1f77bcf86cd799439011'") + ) + row = result.fetchone() + assert row is not None + assert row[0] == "507f1f77bcf86cd799439011" # No prefix in DB diff --git a/tests/unit/test_request_document_translator.py b/tests/unit/test_request_document_translator.py index 93130a89..0ad71c5f 100644 --- a/tests/unit/test_request_document_translator.py +++ b/tests/unit/test_request_document_translator.py @@ -3,8 +3,7 @@ import pytest import pytest_mock - -from askui.chat.api.messages.models import RequestDocumentBlockParam +from askui.chat.api.messages.schemas import RequestDocumentBlockParam from askui.chat.api.messages.translator import RequestDocumentBlockParamTranslator from askui.models.shared.agent_message_param import ( CacheControlEphemeralParam,