chore: fix all ruff lint errors + format — 0 errors, 306 tests pass

This commit is contained in:
leocrm-bot
2026-06-29 17:43:56 +02:00
parent 316f323ff4
commit a2452cc04b
81 changed files with 2317 additions and 1128 deletions
+12 -3
View File
@@ -5,11 +5,11 @@ from __future__ import annotations
import asyncio import asyncio
from logging.config import fileConfig from logging.config import fileConfig
from alembic import context
from sqlalchemy import pool from sqlalchemy import pool
from sqlalchemy.engine import Connection from sqlalchemy.engine import Connection
from sqlalchemy.ext.asyncio import async_engine_from_config from sqlalchemy.ext.asyncio import async_engine_from_config
from alembic import context
from app.config import get_settings from app.config import get_settings
from app.core.db import Base from app.core.db import Base
from app.models import * # noqa: F401,F403 from app.models import * # noqa: F401,F403
@@ -25,7 +25,12 @@ config.set_main_option("sqlalchemy.url", settings.database_url)
def run_migrations_offline() -> None: def run_migrations_offline() -> None:
url = config.get_main_option("sqlalchemy.url") url = config.get_main_option("sqlalchemy.url")
context.configure(url=url, target_metadata=target_metadata, literal_binds=True, dialect_opts={"paramstyle": "named"}) context.configure(
url=url,
target_metadata=target_metadata,
literal_binds=True,
dialect_opts={"paramstyle": "named"},
)
with context.begin_transaction(): with context.begin_transaction():
context.run_migrations() context.run_migrations()
@@ -37,7 +42,11 @@ def do_run_migrations(connection: Connection) -> None:
async def run_async_migrations() -> None: async def run_async_migrations() -> None:
connectable = async_engine_from_config(config.get_section(config.config_ini_section, {}), prefix="sqlalchemy.", poolclass=pool.NullPool) connectable = async_engine_from_config(
config.get_section(config.config_ini_section, {}),
prefix="sqlalchemy.",
poolclass=pool.NullPool,
)
async with connectable.connect() as connection: async with connectable.connect() as connection:
await connection.run_sync(do_run_migrations) await connection.run_sync(do_run_migrations)
await connectable.dispose() await connectable.dispose()
+136 -105
View File
@@ -9,19 +9,26 @@ from __future__ import annotations
import re import re
from typing import Any from typing import Any
# Precompiled patterns for intent detection # Precompiled patterns for intent detection
_PATTERNS = { _PATTERNS = {
'create_company': re.compile(r'\b(create|add|new)\b.*\b(company|firm|organization|organisation)\b', re.IGNORECASE), "create_company": re.compile(
'delete_company': re.compile(r'\b(delete|remove)\b.*\b(company|firm)\b', re.IGNORECASE), r"\b(create|add|new)\b.*\b(company|firm|organization|organisation)\b", re.IGNORECASE
'update_company': re.compile(r'\b(update|edit|modify|change)\b.*\b(company|firm)\b', re.IGNORECASE), ),
'list_company': re.compile(r'\b(list|show|find|search|get|display)\b.*\b(compan|firm)\b', re.IGNORECASE), "delete_company": re.compile(r"\b(delete|remove)\b.*\b(company|firm)\b", re.IGNORECASE),
'list_company2': re.compile(r'\bcompan.*\b(list|all)\b', re.IGNORECASE), "update_company": re.compile(
'create_contact': re.compile(r'\b(create|add|new)\b.*\b(contact|person)\b', re.IGNORECASE), r"\b(update|edit|modify|change)\b.*\b(company|firm)\b", re.IGNORECASE
'list_contact': re.compile(r'\b(list|show|find|search|get|display)\b.*\b(contact|person)\b', re.IGNORECASE), ),
'list_workflow': re.compile(r'\b(list|show|get|display)\b.*\b(workflow)\b', re.IGNORECASE), "list_company": re.compile(
'create_workflow': re.compile(r'\b(create|new|add)\b.*\b(workflow)\b', re.IGNORECASE), r"\b(list|show|find|search|get|display)\b.*\b(compan|firm)\b", re.IGNORECASE
'help': re.compile(r'\b(help|what can you do|assist)\b', re.IGNORECASE), ),
"list_company2": re.compile(r"\bcompan.*\b(list|all)\b", re.IGNORECASE),
"create_contact": re.compile(r"\b(create|add|new)\b.*\b(contact|person)\b", re.IGNORECASE),
"list_contact": re.compile(
r"\b(list|show|find|search|get|display)\b.*\b(contact|person)\b", re.IGNORECASE
),
"list_workflow": re.compile(r"\b(list|show|get|display)\b.*\b(workflow)\b", re.IGNORECASE),
"create_workflow": re.compile(r"\b(create|new|add)\b.*\b(workflow)\b", re.IGNORECASE),
"help": re.compile(r"\b(help|what can you do|assist)\b", re.IGNORECASE),
} }
# Name extraction patterns - using single-quoted strings to avoid escaping issues # Name extraction patterns - using single-quoted strings to avoid escaping issues
@@ -36,10 +43,14 @@ _NAME_PATTERNS = [
# Field extraction patterns # Field extraction patterns
_INDUSTRY_PAT = re.compile(r"industry\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE) _INDUSTRY_PAT = re.compile(r"industry\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE)
_NAME_UPDATE_PAT = re.compile(r"(?:name|rename)\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE) _NAME_UPDATE_PAT = re.compile(
r"(?:name|rename)\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE
)
_PHONE_PAT = re.compile(r"phone\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE) _PHONE_PAT = re.compile(r"phone\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE)
_EMAIL_PAT = re.compile(r"email\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE) _EMAIL_PAT = re.compile(r"email\s+(?:to|:)?\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE)
_SEARCH_PAT = re.compile(r"\b(?:named|called|matching|with name)\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE) _SEARCH_PAT = re.compile(
r"\b(?:named|called|matching|with name)\s+['\"]?([^'\".,]+)['\"]?", re.IGNORECASE
)
def map_query_to_actions(query: str, context: dict[str, Any] | None = None) -> list[dict[str, Any]]: def map_query_to_actions(query: str, context: dict[str, Any] | None = None) -> list[dict[str, Any]]:
@@ -53,109 +64,129 @@ def map_query_to_actions(query: str, context: dict[str, Any] | None = None) -> l
actions: list[dict[str, Any]] = [] actions: list[dict[str, Any]] = []
# --- Company intents --- # --- Company intents ---
if _PATTERNS['create_company'].search(q): if _PATTERNS["create_company"].search(q):
name = _extract_name(query) name = _extract_name(query)
actions.append({ actions.append(
'method': 'POST', {
'path': '/api/v1/companies', "method": "POST",
'body': {'name': name or 'New Company'}, "path": "/api/v1/companies",
'description': f"Create a new company named '{name or 'New Company'}'", "body": {"name": name or "New Company"},
'confidence': 0.9, "description": f"Create a new company named '{name or 'New Company'}'",
}) "confidence": 0.9,
}
)
elif _PATTERNS['delete_company'].search(q): elif _PATTERNS["delete_company"].search(q):
entity_id = context.get('company_id') or context.get('entity_id') entity_id = context.get("company_id") or context.get("entity_id")
if entity_id: if entity_id:
actions.append({ actions.append(
'method': 'DELETE', {
'path': f'/api/v1/companies/{entity_id}', "method": "DELETE",
'body': None, "path": f"/api/v1/companies/{entity_id}",
'description': f'Delete company {entity_id}', "body": None,
'confidence': 0.9, "description": f"Delete company {entity_id}",
}) "confidence": 0.9,
}
)
else: else:
actions.append({ actions.append(
'method': 'DELETE', {
'path': '/api/v1/companies/{id}', "method": "DELETE",
'body': None, "path": "/api/v1/companies/{id}",
'description': 'Delete a company (requires company ID in context or selection)', "body": None,
'confidence': 0.5, "description": "Delete a company (requires company ID in context or selection)",
}) "confidence": 0.5,
}
)
elif _PATTERNS['update_company'].search(q): elif _PATTERNS["update_company"].search(q):
entity_id = context.get('company_id') or context.get('entity_id') entity_id = context.get("company_id") or context.get("entity_id")
path = f'/api/v1/companies/{entity_id}' if entity_id else '/api/v1/companies/{id}' path = f"/api/v1/companies/{entity_id}" if entity_id else "/api/v1/companies/{id}"
actions.append({ actions.append(
'method': 'PATCH', {
'path': path, "method": "PATCH",
'body': _extract_update_fields(query), "path": path,
'description': 'Update company information', "body": _extract_update_fields(query),
'confidence': 0.8, "description": "Update company information",
}) "confidence": 0.8,
}
)
elif _PATTERNS['list_company'].search(q) or _PATTERNS['list_company2'].search(q): elif _PATTERNS["list_company"].search(q) or _PATTERNS["list_company2"].search(q):
search_term = _extract_search_term(query) search_term = _extract_search_term(query)
desc = 'List companies' desc = "List companies"
if search_term: if search_term:
desc += f" matching '{search_term}'" desc += f" matching '{search_term}'"
actions.append({ actions.append(
'method': 'GET', {
'path': '/api/v1/companies', "method": "GET",
'body': None, "path": "/api/v1/companies",
'description': desc, "body": None,
'confidence': 0.85, "description": desc,
}) "confidence": 0.85,
}
)
# --- Contact intents --- # --- Contact intents ---
elif _PATTERNS['create_contact'].search(q): elif _PATTERNS["create_contact"].search(q):
name = _extract_name(query) name = _extract_name(query)
actions.append({ actions.append(
'method': 'POST', {
'path': '/api/v1/contacts', "method": "POST",
'body': {'name': name or 'New Contact'}, "path": "/api/v1/contacts",
'description': f"Create a new contact named '{name or 'New Contact'}'", "body": {"name": name or "New Contact"},
'confidence': 0.9, "description": f"Create a new contact named '{name or 'New Contact'}'",
}) "confidence": 0.9,
}
)
elif _PATTERNS['list_contact'].search(q): elif _PATTERNS["list_contact"].search(q):
actions.append({ actions.append(
'method': 'GET', {
'path': '/api/v1/contacts', "method": "GET",
'body': None, "path": "/api/v1/contacts",
'description': 'List contacts', "body": None,
'confidence': 0.85, "description": "List contacts",
}) "confidence": 0.85,
}
)
# --- Workflow intents --- # --- Workflow intents ---
elif _PATTERNS['list_workflow'].search(q): elif _PATTERNS["list_workflow"].search(q):
actions.append({ actions.append(
'method': 'GET', {
'path': '/api/v1/workflows', "method": "GET",
'body': None, "path": "/api/v1/workflows",
'description': 'List workflows', "body": None,
'confidence': 0.85, "description": "List workflows",
}) "confidence": 0.85,
}
)
elif _PATTERNS['create_workflow'].search(q): elif _PATTERNS["create_workflow"].search(q):
name = _extract_name(query) name = _extract_name(query)
actions.append({ actions.append(
'method': 'POST', {
'path': '/api/v1/workflows', "method": "POST",
'body': {'name': name or 'New Workflow', 'steps': []}, "path": "/api/v1/workflows",
'description': 'Create a new workflow', "body": {"name": name or "New Workflow", "steps": []},
'confidence': 0.8, "description": "Create a new workflow",
}) "confidence": 0.8,
}
)
# --- Generic fallback --- # --- Generic fallback ---
if not actions: if not actions:
if _PATTERNS['help'].search(q): if _PATTERNS["help"].search(q):
actions.append({ actions.append(
'method': 'GET', {
'path': '/api/v1/companies', "method": "GET",
'body': None, "path": "/api/v1/companies",
'description': 'Show available companies (demonstration action)', "body": None,
'confidence': 0.3, "description": "Show available companies (demonstration action)",
}) "confidence": 0.3,
}
)
return actions return actions
@@ -180,20 +211,20 @@ def _extract_search_term(query: str) -> str | None:
def _extract_update_fields(query: str) -> dict[str, Any]: def _extract_update_fields(query: str) -> dict[str, Any]:
"""Extract fields to update from the query.""" """Extract fields to update from the query."""
fields: dict[str, Any] = {} fields: dict[str, Any] = {}
if re.search(r'\bindustry\b', query, re.IGNORECASE): if re.search(r"\bindustry\b", query, re.IGNORECASE):
match = _INDUSTRY_PAT.search(query) match = _INDUSTRY_PAT.search(query)
if match: if match:
fields['industry'] = match.group(1).strip() fields["industry"] = match.group(1).strip()
if re.search(r'\b(name|rename)\b', query, re.IGNORECASE): if re.search(r"\b(name|rename)\b", query, re.IGNORECASE):
match = _NAME_UPDATE_PAT.search(query) match = _NAME_UPDATE_PAT.search(query)
if match: if match:
fields['name'] = match.group(1).strip() fields["name"] = match.group(1).strip()
if re.search(r'\bphone\b', query, re.IGNORECASE): if re.search(r"\bphone\b", query, re.IGNORECASE):
match = _PHONE_PAT.search(query) match = _PHONE_PAT.search(query)
if match: if match:
fields['phone'] = match.group(1).strip() fields["phone"] = match.group(1).strip()
if re.search(r'\bemail\b', query, re.IGNORECASE): if re.search(r"\bemail\b", query, re.IGNORECASE):
match = _EMAIL_PAT.search(query) match = _EMAIL_PAT.search(query)
if match: if match:
fields['email'] = match.group(1).strip() fields["email"] = match.group(1).strip()
return fields or {'name': 'Updated Name'} return fields or {"name": "Updated Name"}
+22 -6
View File
@@ -7,9 +7,9 @@ tests to run without external API dependencies.
from __future__ import annotations from __future__ import annotations
import os
import json import json
import logging import logging
import os
from typing import Any from typing import Any
import httpx import httpx
@@ -20,7 +20,9 @@ logger = logging.getLogger(__name__)
class LLMResponse: class LLMResponse:
"""Structured LLM response containing proposed actions.""" """Structured LLM response containing proposed actions."""
def __init__(self, message: str, proposed_actions: list[dict[str, Any]], confidence: float = 0.8): def __init__(
self, message: str, proposed_actions: list[dict[str, Any]], confidence: float = 0.8
):
self.message = message self.message = message
self.proposed_actions = proposed_actions self.proposed_actions = proposed_actions
self.confidence = confidence self.confidence = confidence
@@ -41,7 +43,9 @@ class LLMClient:
- Otherwise: mock/stub mode with keyword-based action mapping - Otherwise: mock/stub mode with keyword-based action mapping
""" """
def __init__(self, model: str | None = None, api_key: str | None = None, api_base: str | None = None): def __init__(
self, model: str | None = None, api_key: str | None = None, api_base: str | None = None
):
self.model = model or os.environ.get("AI_MODEL", "") self.model = model or os.environ.get("AI_MODEL", "")
self.api_key = api_key or os.environ.get("AI_API_KEY", "") self.api_key = api_key or os.environ.get("AI_API_KEY", "")
self.api_base = api_base or os.environ.get("AI_API_BASE", "https://api.openai.com/v1") self.api_base = api_base or os.environ.get("AI_API_BASE", "https://api.openai.com/v1")
@@ -118,9 +122,21 @@ class LLMClient:
available_apis = [ available_apis = [
{"method": "GET", "path": "/api/v1/companies", "description": "List companies"}, {"method": "GET", "path": "/api/v1/companies", "description": "List companies"},
{"method": "POST", "path": "/api/v1/companies", "description": "Create a company"}, {"method": "POST", "path": "/api/v1/companies", "description": "Create a company"},
{"method": "GET", "path": "/api/v1/companies/{id}", "description": "Get company details"}, {
{"method": "PATCH", "path": "/api/v1/companies/{id}", "description": "Update a company"}, "method": "GET",
{"method": "DELETE", "path": "/api/v1/companies/{id}", "description": "Delete a company"}, "path": "/api/v1/companies/{id}",
"description": "Get company details",
},
{
"method": "PATCH",
"path": "/api/v1/companies/{id}",
"description": "Update a company",
},
{
"method": "DELETE",
"path": "/api/v1/companies/{id}",
"description": "Delete a company",
},
{"method": "GET", "path": "/api/v1/contacts", "description": "List contacts"}, {"method": "GET", "path": "/api/v1/contacts", "description": "List contacts"},
{"method": "POST", "path": "/api/v1/contacts", "description": "Create a contact"}, {"method": "POST", "path": "/api/v1/contacts", "description": "Create a contact"},
{"method": "GET", "path": "/api/v1/workflows", "description": "List workflows"}, {"method": "GET", "path": "/api/v1/workflows", "description": "List workflows"},
+12 -7
View File
@@ -5,20 +5,20 @@ from __future__ import annotations
import hashlib import hashlib
import secrets import secrets
import uuid import uuid
from datetime import datetime, timedelta, timezone from datetime import UTC, datetime, timedelta
from typing import Any from typing import Any
import redis.asyncio as aioredis import redis.asyncio as aioredis
from passlib.context import CryptContext from passlib.context import CryptContext
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings from app.config import get_settings
from app.models.user import User, UserTenant
from app.models.role import Role
from app.models.session import Session as SessionModel from app.models.session import Session as SessionModel
from app.models.user import User
_pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto", bcrypt__rounds=get_settings().bcrypt_rounds) _pwd_context = CryptContext(
schemes=["bcrypt"], deprecated="auto", bcrypt__rounds=get_settings().bcrypt_rounds
)
def hash_password(password: str) -> str: def hash_password(password: str) -> str:
@@ -63,7 +63,7 @@ async def create_session(
settings = get_settings() settings = get_settings()
session_id = str(uuid.uuid4()) session_id = str(uuid.uuid4())
csrf_token = generate_csrf_token() csrf_token = generate_csrf_token()
expires_at = datetime.now(timezone.utc) + timedelta(seconds=settings.session_ttl_seconds) expires_at = datetime.now(UTC) + timedelta(seconds=settings.session_ttl_seconds)
# Redis runtime session # Redis runtime session
session_data: dict[str, Any] = { session_data: dict[str, Any] = {
@@ -76,6 +76,7 @@ async def create_session(
"is_active": user.is_active, "is_active": user.is_active,
} }
import json import json
await redis.setex( await redis.setex(
f"session:{session_id}", f"session:{session_id}",
settings.session_ttl_seconds, settings.session_ttl_seconds,
@@ -99,6 +100,7 @@ async def create_session(
async def get_session_data(redis: aioredis.Redis, session_id: str) -> dict[str, Any] | None: async def get_session_data(redis: aioredis.Redis, session_id: str) -> dict[str, Any] | None:
"""Retrieve session data from Redis.""" """Retrieve session data from Redis."""
import json import json
raw = await redis.get(f"session:{session_id}") raw = await redis.get(f"session:{session_id}")
if raw is None: if raw is None:
return None return None
@@ -117,6 +119,7 @@ async def update_session_tenant(
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""Update the active tenant in a Redis session.""" """Update the active tenant in a Redis session."""
import json import json
settings = get_settings() settings = get_settings()
raw = await redis.get(f"session:{session_id}") raw = await redis.get(f"session:{session_id}")
if raw is None: if raw is None:
@@ -130,7 +133,9 @@ async def update_session_tenant(
return data return data
def check_permission(role_name: str, module: str, action: str, permissions: dict | None = None) -> bool: def check_permission(
role_name: str, module: str, action: str, permissions: dict | None = None
) -> bool:
"""Check if a role has permission for a module+action. """Check if a role has permission for a module+action.
Built-in roles: admin (all), editor (read+write), viewer (read only). Built-in roles: admin (all), editor (read+write), viewer (read only).
Custom roles use the permissions dict. Custom roles use the permissions dict.
-1
View File
@@ -9,7 +9,6 @@ import redis.asyncio as aioredis
from app.config import get_settings from app.config import get_settings
_cache_redis: aioredis.Redis | None = None _cache_redis: aioredis.Redis | None = None
+13 -15
View File
@@ -3,10 +3,14 @@
from __future__ import annotations from __future__ import annotations
import contextlib import contextlib
import uuid
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from typing import Any from datetime import datetime
from typing import Any # noqa: F401
from sqlalchemy import event as sa_event from sqlalchemy import DateTime, String, func, text # noqa: F401
from sqlalchemy import event as sa_event # noqa: F401
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.ext.asyncio import ( from sqlalchemy.ext.asyncio import (
AsyncEngine, AsyncEngine,
AsyncSession, AsyncSession,
@@ -14,10 +18,6 @@ from sqlalchemy.ext.asyncio import (
create_async_engine, create_async_engine,
) )
from sqlalchemy.orm import DeclarativeBase, Mapped, declared_attr, mapped_column from sqlalchemy.orm import DeclarativeBase, Mapped, declared_attr, mapped_column
from sqlalchemy import text, String, DateTime, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID
import uuid
from datetime import datetime
from app.config import get_settings from app.config import get_settings
@@ -26,8 +26,8 @@ class Base(DeclarativeBase):
"""Declarative base with shared columns and tenant-scoping support.""" """Declarative base with shared columns and tenant-scoping support."""
@declared_attr.directive @declared_attr.directive
def __tablename__(cls) -> str: def __tablename__(self) -> str:
return cls.__name__.lower() + "s" return self.__name__.lower() + "s"
class TimestampMixin: class TimestampMixin:
@@ -44,9 +44,7 @@ class TimestampMixin:
class TenantMixin(TimestampMixin): class TenantMixin(TimestampMixin):
"""Adds tenant_id column and enables ORM-level auto-filtering.""" """Adds tenant_id column and enables ORM-level auto-filtering."""
tenant_id: Mapped[uuid.UUID] = mapped_column( tenant_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False, index=True)
PGUUID(as_uuid=True), nullable=False, index=True
)
# Global engine and session factory # Global engine and session factory
@@ -101,7 +99,9 @@ async def set_tenant_context(session: AsyncSession, tenant_id: uuid.UUID | str)
@contextlib.asynccontextmanager @contextlib.asynccontextmanager
async def create_db_session(tenant_id: uuid.UUID | str | None = None) -> AsyncGenerator[AsyncSession, None]: async def create_db_session(
tenant_id: uuid.UUID | str | None = None,
) -> AsyncGenerator[AsyncSession, None]:
"""Create a standalone session outside FastAPI (e.g. for tests/workers).""" """Create a standalone session outside FastAPI (e.g. for tests/workers)."""
factory = get_session_factory() factory = get_session_factory()
async with factory() as session: async with factory() as session:
@@ -128,7 +128,5 @@ def reset_engine_for_testing(engine: AsyncEngine) -> async_sessionmaker[AsyncSes
"""Replace the global engine with a test engine. Returns a session factory.""" """Replace the global engine with a test engine. Returns a session factory."""
global _engine, _session_factory global _engine, _session_factory
_engine = engine _engine = engine
_session_factory = async_sessionmaker( _session_factory = async_sessionmaker(bind=engine, expire_on_commit=False, class_=AsyncSession)
bind=engine, expire_on_commit=False, class_=AsyncSession
)
return _session_factory return _session_factory
+3 -1
View File
@@ -4,7 +4,8 @@ from __future__ import annotations
import asyncio import asyncio
from collections import defaultdict from collections import defaultdict
from typing import Any, Callable, Coroutine from collections.abc import Callable, Coroutine
from typing import Any
EventHandler = Callable[[dict[str, Any]], Coroutine[Any, Any, None]] EventHandler = Callable[[dict[str, Any]], Coroutine[Any, Any, None]]
@@ -48,4 +49,5 @@ def register_workflow_event_handlers() -> None:
Should be called during application startup. Should be called during application startup.
""" """
from app.workflows.engine import register_workflow_event_handlers as _register from app.workflows.engine import register_workflow_event_handlers as _register
_register() _register()
+20 -10
View File
@@ -3,9 +3,10 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import UTC
from typing import Any from typing import Any
from sqlalchemy import select, func, update from sqlalchemy import func, select, update
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.models.notification import Notification from app.models.notification import Notification
@@ -42,9 +43,13 @@ async def list_notifications(
"""List notifications for a user, unread first, then by created_at desc.""" """List notifications for a user, unread first, then by created_at desc."""
offset = (page - 1) * page_size offset = (page - 1) * page_size
# Count total # Count total
count_q = select(func.count()).select_from(Notification).where( count_q = (
Notification.tenant_id == tenant_id, select(func.count())
Notification.user_id == user_id, .select_from(Notification)
.where(
Notification.tenant_id == tenant_id,
Notification.user_id == user_id,
)
) )
total = (await db.execute(count_q)).scalar() or 0 total = (await db.execute(count_q)).scalar() or 0
@@ -80,7 +85,8 @@ async def mark_notification_read(
notification_id: uuid.UUID, notification_id: uuid.UUID,
) -> Notification | None: ) -> Notification | None:
"""Mark a notification as read.""" """Mark a notification as read."""
from datetime import datetime, timezone from datetime import datetime
q = ( q = (
update(Notification) update(Notification)
.where( .where(
@@ -88,7 +94,7 @@ async def mark_notification_read(
Notification.tenant_id == tenant_id, Notification.tenant_id == tenant_id,
Notification.user_id == user_id, Notification.user_id == user_id,
) )
.values(read_at=datetime.now(timezone.utc)) .values(read_at=datetime.now(UTC))
.returning(Notification) .returning(Notification)
) )
result = await db.execute(q) result = await db.execute(q)
@@ -102,10 +108,14 @@ async def get_unread_count(
user_id: uuid.UUID, user_id: uuid.UUID,
) -> int: ) -> int:
"""Get unread notification count for a user.""" """Get unread notification count for a user."""
q = select(func.count()).select_from(Notification).where( q = (
Notification.tenant_id == tenant_id, select(func.count())
Notification.user_id == user_id, .select_from(Notification)
Notification.read_at.is_(None), .where(
Notification.tenant_id == tenant_id,
Notification.user_id == user_id,
Notification.read_at.is_(None),
)
) )
result = await db.execute(q) result = await db.execute(q)
return result.scalar() or 0 return result.scalar() or 0
-1
View File
@@ -5,7 +5,6 @@ from __future__ import annotations
from fastapi import HTTPException, Request, status from fastapi import HTTPException, Request, status
from app.core.auth import get_redis from app.core.auth import get_redis
from app.config import get_settings
async def check_rate_limit( async def check_rate_limit(
+1 -3
View File
@@ -3,10 +3,7 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from typing import Any
from sqlalchemy import event as sa_event
from sqlalchemy.orm import Session, with_loader_criteria
from sqlalchemy.sql import Select from sqlalchemy.sql import Select
from app.core.db import TenantMixin from app.core.db import TenantMixin
@@ -21,4 +18,5 @@ def apply_tenant_filter(query: Select, tenant_id: uuid.UUID) -> Select:
async def set_rls_context(session, tenant_id: uuid.UUID | str) -> None: async def set_rls_context(session, tenant_id: uuid.UUID | str) -> None:
"""Set PostgreSQL RLS session variable.""" """Set PostgreSQL RLS session variable."""
from app.core.db import set_tenant_context from app.core.db import set_tenant_context
await set_tenant_context(session, tenant_id) await set_tenant_context(session, tenant_id)
-2
View File
@@ -7,13 +7,11 @@ from typing import Any
import redis.asyncio as aioredis import redis.asyncio as aioredis
from fastapi import Depends, HTTPException, Request, status from fastapi import Depends, HTTPException, Request, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings from app.config import get_settings
from app.core.auth import get_redis, get_session_data from app.core.auth import get_redis, get_session_data
from app.core.db import get_db, set_tenant_context from app.core.db import get_db, set_tenant_context
from app.models.user import User
async def get_redis_dep() -> aioredis.Redis: async def get_redis_dep() -> aioredis.Redis:
+15 -2
View File
@@ -8,11 +8,24 @@ from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from app.config import get_settings from app.config import get_settings
from app.core.middleware import CSRFMiddleware
from app.core.db import close_engine, get_engine from app.core.db import close_engine, get_engine
from app.core.middleware import CSRFMiddleware
from app.core.service_container import get_container from app.core.service_container import get_container
from app.plugins.registry import get_registry from app.plugins.registry import get_registry
from app.routes import auth, users, roles, tenants, health, notifications, companies, contacts, import_export, plugins, ai_copilot, workflows from app.routes import (
ai_copilot,
auth,
companies,
contacts,
health,
import_export,
notifications,
plugins,
roles,
tenants,
users,
workflows,
)
@asynccontextmanager @asynccontextmanager
+43 -20
View File
@@ -1,32 +1,55 @@
"""SQLAlchemy models for LeoCRM.""" """SQLAlchemy models for LeoCRM."""
from app.models.tenant import Tenant from app.models.audit import AuditLog, DeletionLog
from app.models.user import User, UserTenant from app.models.auth import ApiToken, PasswordResetToken
from app.models.company import Company
from app.models.contact import CompanyContact, Contact
from app.models.notification import Notification
from app.models.plugin import Plugin, PluginMigration
from app.models.role import Role from app.models.role import Role
from app.models.session import Session from app.models.session import Session
from app.models.audit import AuditLog, DeletionLog from app.models.tenant import Tenant
from app.models.notification import Notification from app.models.user import User, UserTenant
from app.models.auth import PasswordResetToken, ApiToken
from app.models.company import Company
from app.models.contact import Contact, CompanyContact
from app.models.plugin import Plugin, PluginMigration
__all__ = [ __all__ = [
"Tenant", "User", "UserTenant", "Role", "Session", "Tenant",
"AuditLog", "DeletionLog", "Notification", "User",
"PasswordResetToken", "ApiToken", "Company", "UserTenant",
"Contact", "CompanyContact", "Role",
"Plugin", "PluginMigration", "Session",
"AuditLog",
"DeletionLog",
"Notification",
"PasswordResetToken",
"ApiToken",
"Company",
"Contact",
"CompanyContact",
"Plugin",
"PluginMigration",
] ]
from app.models.ai_conversation import AIConversation, AIMessage from app.models.ai_conversation import AIConversation, AIMessage
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
__all__ = [ __all__ = [
"Tenant", "User", "UserTenant", "Role", "Session", "Tenant",
"AuditLog", "DeletionLog", "Notification", "User",
"PasswordResetToken", "ApiToken", "Company", "UserTenant",
"Contact", "CompanyContact", "Role",
"Plugin", "PluginMigration", "Session",
"AIConversation", "AIMessage", "AuditLog",
"Workflow", "WorkflowInstance", "WorkflowStepHistory", "DeletionLog",
"Notification",
"PasswordResetToken",
"ApiToken",
"Company",
"Contact",
"CompanyContact",
"Plugin",
"PluginMigration",
"AIConversation",
"AIMessage",
"Workflow",
"WorkflowInstance",
"WorkflowStepHistory",
] ]
+9 -10
View File
@@ -3,11 +3,11 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, Text, DateTime, ForeignKey, func, Index, Integer from sqlalchemy import ForeignKey, Index, Integer, String, Text
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
@@ -17,9 +17,7 @@ class AIConversation(Base, TenantMixin):
"""AI Copilot conversation thread — tenant-scoped.""" """AI Copilot conversation thread — tenant-scoped."""
__tablename__ = "ai_conversations" __tablename__ = "ai_conversations"
__table_args__ = ( __table_args__ = (Index("ix_ai_conversations_tenant_user", "tenant_id", "user_id"),)
Index("ix_ai_conversations_tenant_user", "tenant_id", "user_id"),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
@@ -35,15 +33,16 @@ class AIMessage(Base, TenantMixin):
"""Individual messages within an AI conversation — user input, AI response, actions.""" """Individual messages within an AI conversation — user input, AI response, actions."""
__tablename__ = "ai_messages" __tablename__ = "ai_messages"
__table_args__ = ( __table_args__ = (Index("ix_ai_messages_tenant_conversation", "tenant_id", "conversation_id"),)
Index("ix_ai_messages_tenant_conversation", "tenant_id", "conversation_id"),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
) )
conversation_id: Mapped[uuid.UUID] = mapped_column( conversation_id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("ai_conversations.id", ondelete="CASCADE"), nullable=False, index=True PGUUID(as_uuid=True),
ForeignKey("ai_conversations.id", ondelete="CASCADE"),
nullable=False,
index=True,
) )
role: Mapped[str] = mapped_column(String(20), nullable=False) # user, assistant, system role: Mapped[str] = mapped_column(String(20), nullable=False) # user, assistant, system
content: Mapped[str] = mapped_column(Text, nullable=False) content: Mapped[str] = mapped_column(Text, nullable=False)
+3 -2
View File
@@ -6,8 +6,9 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, DateTime, ForeignKey, func, Text from sqlalchemy import DateTime, ForeignKey, String, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
+4 -5
View File
@@ -5,8 +5,9 @@ from __future__ import annotations
import uuid import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import String, DateTime, ForeignKey, func, Index from sqlalchemy import DateTime, ForeignKey, Index, String, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
@@ -32,9 +33,7 @@ class ApiToken(Base, TenantMixin):
"""API token for programmatic access.""" """API token for programmatic access."""
__tablename__ = "api_tokens" __tablename__ = "api_tokens"
__table_args__ = ( __table_args__ = (Index("ix_api_tokens_tenant_user", "tenant_id", "user_id"),)
Index("ix_api_tokens_tenant_user", "tenant_id", "user_id"),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
+4 -5
View File
@@ -6,8 +6,9 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, Text, DateTime, ForeignKey, Index, Computed from sqlalchemy import Computed, DateTime, ForeignKey, Index, String, Text
from sqlalchemy.dialects.postgresql import UUID as PGUUID, TSVECTOR from sqlalchemy.dialects.postgresql import TSVECTOR
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
@@ -34,9 +35,7 @@ class Company(Base, TenantMixin):
email: Mapped[str | None] = mapped_column(String(255), nullable=True) email: Mapped[str | None] = mapped_column(String(255), nullable=True)
website: Mapped[str | None] = mapped_column(String(500), nullable=True) website: Mapped[str | None] = mapped_column(String(500), nullable=True)
description: Mapped[str | None] = mapped_column(Text, nullable=True) description: Mapped[str | None] = mapped_column(Text, nullable=True)
deleted_at: Mapped[datetime | None] = mapped_column( deleted_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
DateTime(timezone=True), nullable=True
)
# FTS vector — PostgreSQL generated column, auto-updated on insert/update # FTS vector — PostgreSQL generated column, auto-updated on insert/update
search_tsv: Mapped[Any] = mapped_column( search_tsv: Mapped[Any] = mapped_column(
TSVECTOR, TSVECTOR,
+5 -9
View File
@@ -6,12 +6,12 @@ import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import ( from sqlalchemy import (
String,
Text,
DateTime,
Boolean, Boolean,
DateTime,
ForeignKey, ForeignKey,
Index, Index,
String,
Text,
UniqueConstraint, UniqueConstraint,
) )
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
@@ -42,9 +42,7 @@ class Contact(Base, TenantMixin):
department: Mapped[str | None] = mapped_column(String(100), nullable=True) department: Mapped[str | None] = mapped_column(String(100), nullable=True)
linkedin_url: Mapped[str | None] = mapped_column(String(500), nullable=True) linkedin_url: Mapped[str | None] = mapped_column(String(500), nullable=True)
notes: Mapped[str | None] = mapped_column(Text, nullable=True) notes: Mapped[str | None] = mapped_column(Text, nullable=True)
deleted_at: Mapped[datetime | None] = mapped_column( deleted_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
DateTime(timezone=True), nullable=True
)
created_by: Mapped[uuid.UUID | None] = mapped_column( created_by: Mapped[uuid.UUID | None] = mapped_column(
PGUUID(as_uuid=True), PGUUID(as_uuid=True),
ForeignKey("users.id", ondelete="SET NULL"), ForeignKey("users.id", ondelete="SET NULL"),
@@ -62,9 +60,7 @@ class CompanyContact(Base, TenantMixin):
__tablename__ = "company_contacts" __tablename__ = "company_contacts"
__table_args__ = ( __table_args__ = (
UniqueConstraint( UniqueConstraint("company_id", "contact_id", "tenant_id", name="uq_company_contact_tenant"),
"company_id", "contact_id", "tenant_id", name="uq_company_contact_tenant"
),
Index("ix_cc_company", "company_id"), Index("ix_cc_company", "company_id"),
Index("ix_cc_contact", "contact_id"), Index("ix_cc_contact", "contact_id"),
) )
+1 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import String, DateTime, ForeignKey, Text, func, Index from sqlalchemy import DateTime, ForeignKey, Index, String, Text, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
+8 -18
View File
@@ -3,10 +3,9 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, Boolean, Text, DateTime, func, Index from sqlalchemy import Boolean, Index, String, Text
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
@@ -21,9 +20,7 @@ class Plugin(Base, TimestampMixin):
""" """
__tablename__ = "plugins" __tablename__ = "plugins"
__table_args__ = ( __table_args__ = (Index("ix_plugins_name", "name", unique=True),)
Index("ix_plugins_name", "name", unique=True),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
@@ -31,17 +28,12 @@ class Plugin(Base, TimestampMixin):
name: Mapped[str] = mapped_column(String(80), nullable=False, unique=True) name: Mapped[str] = mapped_column(String(80), nullable=False, unique=True)
display_name: Mapped[str] = mapped_column(String(120), nullable=False) display_name: Mapped[str] = mapped_column(String(120), nullable=False)
version: Mapped[str] = mapped_column(String(40), nullable=False) version: Mapped[str] = mapped_column(String(40), nullable=False)
status: Mapped[str] = mapped_column( status: Mapped[str] = mapped_column(String(20), nullable=False, default="installed")
String(20), nullable=False, default="installed" installed: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True)
) active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
installed: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=True
)
active: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=False
)
config: Mapped[dict[str, Any] | None] = mapped_column( config: Mapped[dict[str, Any] | None] = mapped_column(
Text, nullable=True # JSON string for plugin configuration Text,
nullable=True, # JSON string for plugin configuration
) )
# Transient attribute for response (not persisted) # Transient attribute for response (not persisted)
@@ -63,6 +55,4 @@ class PluginMigration(Base, TimestampMixin):
) )
plugin_name: Mapped[str] = mapped_column(String(80), nullable=False) plugin_name: Mapped[str] = mapped_column(String(80), nullable=False)
migration_file: Mapped[str] = mapped_column(String(255), nullable=False) migration_file: Mapped[str] = mapped_column(String(255), nullable=False)
status: Mapped[str] = mapped_column( status: Mapped[str] = mapped_column(String(20), nullable=False, default="applied")
String(20), nullable=False, default="applied"
)
+3 -2
View File
@@ -6,8 +6,9 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, DateTime, ForeignKey, func from sqlalchemy import DateTime, String, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
+1 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import String, DateTime, ForeignKey, func from sqlalchemy import DateTime, ForeignKey, String, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
+1 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import String, DateTime, func from sqlalchemy import DateTime, String, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
+5 -6
View File
@@ -6,9 +6,10 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, Boolean, DateTime, ForeignKey, func, UniqueConstraint from sqlalchemy import Boolean, DateTime, ForeignKey, String, UniqueConstraint, func
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
@@ -17,9 +18,7 @@ class User(Base, TenantMixin):
"""User entity — belongs to a tenant, can be member of multiple tenants.""" """User entity — belongs to a tenant, can be member of multiple tenants."""
__tablename__ = "users" __tablename__ = "users"
__table_args__ = ( __table_args__ = (UniqueConstraint("tenant_id", "email", name="uq_users_tenant_email"),)
UniqueConstraint("tenant_id", "email", name="uq_users_tenant_email"),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
+17 -14
View File
@@ -6,8 +6,9 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import String, Text, DateTime, ForeignKey, func, Index, Integer, Boolean from sqlalchemy import Boolean, DateTime, ForeignKey, Index, Integer, String, Text
from sqlalchemy.dialects.postgresql import UUID as PGUUID, JSONB from sqlalchemy.dialects.postgresql import JSONB
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from app.core.db import Base, TenantMixin from app.core.db import Base, TenantMixin
@@ -48,7 +49,10 @@ class WorkflowInstance(Base, TenantMixin):
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
) )
workflow_id: Mapped[uuid.UUID] = mapped_column( workflow_id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("workflows.id", ondelete="CASCADE"), nullable=False, index=True PGUUID(as_uuid=True),
ForeignKey("workflows.id", ondelete="CASCADE"),
nullable=False,
index=True,
) )
status: Mapped[str] = mapped_column(String(30), nullable=False, default="pending") status: Mapped[str] = mapped_column(String(30), nullable=False, default="pending")
# pending → in_progress → completed / rejected / cancelled # pending → in_progress → completed / rejected / cancelled
@@ -57,32 +61,31 @@ class WorkflowInstance(Base, TenantMixin):
initiated_by: Mapped[uuid.UUID | None] = mapped_column( initiated_by: Mapped[uuid.UUID | None] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("users.id", ondelete="SET NULL"), nullable=True PGUUID(as_uuid=True), ForeignKey("users.id", ondelete="SET NULL"), nullable=True
) )
completed_at: Mapped[datetime | None] = mapped_column( completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
DateTime(timezone=True), nullable=True
)
timeout_hours: Mapped[int | None] = mapped_column(Integer, nullable=True) timeout_hours: Mapped[int | None] = mapped_column(Integer, nullable=True)
timeout_at: Mapped[datetime | None] = mapped_column( timeout_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
DateTime(timezone=True), nullable=True
)
class WorkflowStepHistory(Base, TenantMixin): class WorkflowStepHistory(Base, TenantMixin):
"""Immutable record of every step transition in a workflow instance.""" """Immutable record of every step transition in a workflow instance."""
__tablename__ = "workflow_step_history" __tablename__ = "workflow_step_history"
__table_args__ = ( __table_args__ = (Index("ix_wf_step_history_tenant_instance", "tenant_id", "instance_id"),)
Index("ix_wf_step_history_tenant_instance", "tenant_id", "instance_id"),
)
id: Mapped[uuid.UUID] = mapped_column( id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4 PGUUID(as_uuid=True), primary_key=True, default=uuid.uuid4
) )
instance_id: Mapped[uuid.UUID] = mapped_column( instance_id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("workflow_instances.id", ondelete="CASCADE"), nullable=False, index=True PGUUID(as_uuid=True),
ForeignKey("workflow_instances.id", ondelete="CASCADE"),
nullable=False,
index=True,
) )
step_index: Mapped[int] = mapped_column(Integer, nullable=False) step_index: Mapped[int] = mapped_column(Integer, nullable=False)
step_type: Mapped[str] = mapped_column(String(50), nullable=False) step_type: Mapped[str] = mapped_column(String(50), nullable=False)
action: Mapped[str] = mapped_column(String(50), nullable=False) # entered, approved, rejected, skipped, cancelled action: Mapped[str] = mapped_column(
String(50), nullable=False
) # entered, approved, rejected, skipped, cancelled
actor_id: Mapped[uuid.UUID | None] = mapped_column( actor_id: Mapped[uuid.UUID | None] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("users.id", ondelete="SET NULL"), nullable=True PGUUID(as_uuid=True), ForeignKey("users.id", ondelete="SET NULL"), nullable=True
) )
+2 -2
View File
@@ -1,8 +1,8 @@
"""Plugin system framework for LeoCRM.""" """Plugin system framework for LeoCRM."""
from app.plugins.manifest import PluginManifest
from app.plugins.base import BasePlugin from app.plugins.base import BasePlugin
from app.plugins.registry import get_registry from app.plugins.manifest import PluginManifest
from app.plugins.migration_runner import MigrationRunner from app.plugins.migration_runner import MigrationRunner
from app.plugins.registry import get_registry
__all__ = ["PluginManifest", "BasePlugin", "get_registry", "MigrationRunner"] __all__ = ["PluginManifest", "BasePlugin", "get_registry", "MigrationRunner"]
+7 -1
View File
@@ -11,6 +11,7 @@ from app.plugins.manifest import PluginManifest
if TYPE_CHECKING: if TYPE_CHECKING:
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.event_bus import EventBus from app.core.event_bus import EventBus
from app.core.service_container import ServiceContainer from app.core.service_container import ServiceContainer
@@ -85,6 +86,7 @@ class BasePlugin(ABC):
routers: list[APIRouter] = [] routers: list[APIRouter] = []
for route_def in self.manifest.routes: for route_def in self.manifest.routes:
import importlib import importlib
module = importlib.import_module(route_def.module) module = importlib.import_module(route_def.module)
router: APIRouter = getattr(module, route_def.router_attr) router: APIRouter = getattr(module, route_def.router_attr)
routers.append(router) routers.append(router)
@@ -106,9 +108,11 @@ class BasePlugin(ABC):
fallback = getattr(self, "on_event", None) fallback = getattr(self, "on_event", None)
if fallback is not None: if fallback is not None:
return fallback return fallback
# Default no-op handler # Default no-op handler
async def _noop_handler(payload: dict[str, Any]) -> None: async def _noop_handler(payload: dict[str, Any]) -> None:
pass pass
return _noop_handler return _noop_handler
# ─── Utility ─── # ─── Utility ───
@@ -122,4 +126,6 @@ class BasePlugin(ABC):
return self.manifest.version return self.manifest.version
def __repr__(self) -> str: def __repr__(self) -> str:
return f"<{self.__class__.__name__} name={self.manifest.name} version={self.manifest.version}>" return (
f"<{self.__class__.__name__} name={self.manifest.name} version={self.manifest.version}>"
)
+2 -2
View File
@@ -7,8 +7,8 @@ Subdirectory plugins (tags, permissions, entity_links) export their plugin
class via __init__.py so the registry can discover them as packages. class via __init__.py so the registry can discover them as packages.
""" """
from app.plugins.builtins.tags import TagsPlugin
from app.plugins.builtins.permissions import PermissionsPlugin
from app.plugins.builtins.entity_links import EntityLinksPlugin from app.plugins.builtins.entity_links import EntityLinksPlugin
from app.plugins.builtins.permissions import PermissionsPlugin
from app.plugins.builtins.tags import TagsPlugin
__all__ = ["TagsPlugin", "PermissionsPlugin", "EntityLinksPlugin"] __all__ = ["TagsPlugin", "PermissionsPlugin", "EntityLinksPlugin"]
+6 -5
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import uuid import uuid
from sqlalchemy import String, ForeignKey, Index, UniqueConstraint from sqlalchemy import Index, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
@@ -17,7 +17,10 @@ class EntityLink(Base, TenantMixin):
__tablename__ = "entity_links" __tablename__ = "entity_links"
__table_args__ = ( __table_args__ = (
UniqueConstraint( UniqueConstraint(
"tenant_id", "file_id", "entity_type", "entity_id", "tenant_id",
"file_id",
"entity_type",
"entity_id",
name="uq_entity_links_file_entity", name="uq_entity_links_file_entity",
), ),
Index("ix_entity_links_file", "tenant_id", "file_id"), Index("ix_entity_links_file", "tenant_id", "file_id"),
@@ -30,6 +33,4 @@ class EntityLink(Base, TenantMixin):
file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False) file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False)
entity_type: Mapped[str] = mapped_column(String(20), nullable=False) entity_type: Mapped[str] = mapped_column(String(20), nullable=False)
entity_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False) entity_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False)
created_by: Mapped[uuid.UUID | None] = mapped_column( created_by: Mapped[uuid.UUID | None] = mapped_column(PGUUID(as_uuid=True), nullable=True)
PGUUID(as_uuid=True), nullable=True
)
@@ -42,6 +42,7 @@ class EntityLinksPlugin(BasePlugin):
async def on_company_deleted(self, payload: dict[str, Any]) -> None: async def on_company_deleted(self, payload: dict[str, Any]) -> None:
"""Handle company.deleted event — remove all EntityLink rows for that company.""" """Handle company.deleted event — remove all EntityLink rows for that company."""
from sqlalchemy import delete from sqlalchemy import delete
from app.core.db import get_session_factory from app.core.db import get_session_factory
from app.plugins.builtins.entity_links.models import EntityLink from app.plugins.builtins.entity_links.models import EntityLink
@@ -51,6 +52,7 @@ class EntityLinksPlugin(BasePlugin):
return return
import uuid as _uuid import uuid as _uuid
factory = get_session_factory() factory = get_session_factory()
async with factory() as session: async with factory() as session:
await session.execute( await session.execute(
@@ -65,6 +67,7 @@ class EntityLinksPlugin(BasePlugin):
async def on_contact_deleted(self, payload: dict[str, Any]) -> None: async def on_contact_deleted(self, payload: dict[str, Any]) -> None:
"""Handle contact.deleted event — remove all EntityLink rows for that contact.""" """Handle contact.deleted event — remove all EntityLink rows for that contact."""
from sqlalchemy import delete from sqlalchemy import delete
from app.core.db import get_session_factory from app.core.db import get_session_factory
from app.plugins.builtins.entity_links.models import EntityLink from app.plugins.builtins.entity_links.models import EntityLink
@@ -74,6 +77,7 @@ class EntityLinksPlugin(BasePlugin):
return return
import uuid as _uuid import uuid as _uuid
factory = get_session_factory() factory = get_session_factory()
async with factory() as session: async with factory() as session:
await session.execute( await session.execute(
+7 -3
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from fastapi import APIRouter, Body, Depends, HTTPException, Response, status from fastapi import APIRouter, Body, Depends, HTTPException, Response, status
from sqlalchemy import select, delete from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.db import get_db from app.core.db import get_db
@@ -24,7 +24,9 @@ def _parse_uuid(val: str, field: str) -> uuid.UUID:
try: try:
return uuid.UUID(val) return uuid.UUID(val)
except (ValueError, TypeError): except (ValueError, TypeError):
raise HTTPException(400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}
) from None
@router.post("/files/{file_id}/link") @router.post("/files/{file_id}/link")
@@ -41,7 +43,9 @@ async def link_file_to_entity(
entity_id = _parse_uuid(body.entity_id, "entity_id") entity_id = _parse_uuid(body.entity_id, "entity_id")
if body.entity_type not in VALID_ENTITY_TYPES: if body.entity_type not in VALID_ENTITY_TYPES:
raise HTTPException(400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}) raise HTTPException(
400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}
)
# Check if link already exists # Check if link already exists
existing = await db.execute( existing = await db.execute(
+8 -11
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from datetime import datetime from datetime import datetime
from sqlalchemy import String, DateTime, ForeignKey, Index, UniqueConstraint from sqlalchemy import DateTime, Index, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
@@ -18,7 +18,10 @@ class Permission(Base, TenantMixin):
__tablename__ = "permissions" __tablename__ = "permissions"
__table_args__ = ( __table_args__ = (
UniqueConstraint( UniqueConstraint(
"tenant_id", "file_id", "user_id", "access_level", "tenant_id",
"file_id",
"user_id",
"access_level",
name="uq_permissions_file_user_level", name="uq_permissions_file_user_level",
), ),
Index("ix_permissions_file", "tenant_id", "file_id"), Index("ix_permissions_file", "tenant_id", "file_id"),
@@ -30,9 +33,7 @@ class Permission(Base, TenantMixin):
) )
file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False) file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False)
user_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False) user_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False)
group_id: Mapped[uuid.UUID | None] = mapped_column( group_id: Mapped[uuid.UUID | None] = mapped_column(PGUUID(as_uuid=True), nullable=True)
PGUUID(as_uuid=True), nullable=True
)
access_level: Mapped[str] = mapped_column(String(10), nullable=False, default="read") access_level: Mapped[str] = mapped_column(String(10), nullable=False, default="read")
@@ -51,10 +52,6 @@ class ShareLink(Base, TenantMixin):
file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False) file_id: Mapped[uuid.UUID] = mapped_column(PGUUID(as_uuid=True), nullable=False)
token: Mapped[str] = mapped_column(String(64), nullable=False, unique=True) token: Mapped[str] = mapped_column(String(64), nullable=False, unique=True)
password_hash: Mapped[str | None] = mapped_column(String(255), nullable=True) password_hash: Mapped[str | None] = mapped_column(String(255), nullable=True)
expires_at: Mapped[datetime | None] = mapped_column( expires_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
DateTime(timezone=True), nullable=True
)
access_level: Mapped[str] = mapped_column(String(10), nullable=False, default="download") access_level: Mapped[str] = mapped_column(String(10), nullable=False, default="download")
created_by: Mapped[uuid.UUID | None] = mapped_column( created_by: Mapped[uuid.UUID | None] = mapped_column(PGUUID(as_uuid=True), nullable=True)
PGUUID(as_uuid=True), nullable=True
)
@@ -2,10 +2,6 @@
from __future__ import annotations from __future__ import annotations
from typing import Any
from fastapi import APIRouter
from app.plugins.base import BasePlugin from app.plugins.base import BasePlugin
from app.plugins.manifest import PluginManifest, PluginRouteDef from app.plugins.manifest import PluginManifest, PluginRouteDef
+27 -16
View File
@@ -4,10 +4,10 @@ from __future__ import annotations
import secrets import secrets
import uuid import uuid
from datetime import datetime, timezone from datetime import UTC, datetime
from fastapi import APIRouter, Depends, HTTPException, Response, status from fastapi import APIRouter, Depends, HTTPException, Response, status
from sqlalchemy import select, delete from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.auth import hash_password, verify_password from app.core.auth import hash_password, verify_password
@@ -28,11 +28,14 @@ def _parse_uuid(val: str, field: str) -> uuid.UUID:
try: try:
return uuid.UUID(val) return uuid.UUID(val)
except (ValueError, TypeError): except (ValueError, TypeError):
raise HTTPException(400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}
) from None
# ─── Permissions CRUD ─── # ─── Permissions CRUD ───
@router.get("/files/{file_id}/permissions") @router.get("/files/{file_id}/permissions")
async def list_permissions( async def list_permissions(
file_id: str, file_id: str,
@@ -85,7 +88,9 @@ async def grant_permission(
) )
) )
if existing.scalar_one_or_none() is not None: if existing.scalar_one_or_none() is not None:
raise HTTPException(409, detail={"detail": "Permission already exists", "code": "duplicate"}) raise HTTPException(
409, detail={"detail": "Permission already exists", "code": "duplicate"}
)
perm = Permission( perm = Permission(
tenant_id=tenant_id, tenant_id=tenant_id,
@@ -135,8 +140,13 @@ async def revoke_permission(
# ─── Permission Check Helper ─── # ─── Permission Check Helper ───
async def check_user_file_permission( async def check_user_file_permission(
db: AsyncSession, tenant_id: uuid.UUID, file_id: uuid.UUID, user_id: uuid.UUID, required: str = "read" db: AsyncSession,
tenant_id: uuid.UUID,
file_id: uuid.UUID,
user_id: uuid.UUID,
required: str = "read",
) -> bool: ) -> bool:
"""Check if user has required permission on a file.""" """Check if user has required permission on a file."""
result = await db.execute( result = await db.execute(
@@ -158,6 +168,7 @@ async def check_user_file_permission(
# ─── Share Links ─── # ─── Share Links ───
@router.post("/files/{file_id}/share-link") @router.post("/files/{file_id}/share-link")
async def create_share_link( async def create_share_link(
file_id: str, file_id: str,
@@ -219,6 +230,7 @@ async def revoke_share_link(
# ─── Public Share Access (NO AUTH) ─── # ─── Public Share Access (NO AUTH) ───
@public_router.get("/share/{token}") @public_router.get("/share/{token}")
async def public_access( async def public_access(
token: str, token: str,
@@ -229,23 +241,20 @@ async def public_access(
Returns 410 Gone if link is expired. Returns 410 Gone if link is expired.
Returns 403 if password is required but not provided/invalid. Returns 403 if password is required but not provided/invalid.
""" """
result = await db.execute( result = await db.execute(select(ShareLink).where(ShareLink.token == token))
select(ShareLink).where(ShareLink.token == token)
)
link = result.scalar_one_or_none() link = result.scalar_one_or_none()
if link is None: if link is None:
raise HTTPException(404, detail={"detail": "Share link not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Share link not found", "code": "not_found"})
# Check expiry # Check expiry
if link.expires_at is not None: if link.expires_at is not None:
now = datetime.now(timezone.utc) now = datetime.now(UTC)
if now > link.expires_at: if now > link.expires_at:
raise HTTPException(410, detail={"detail": "Share link expired", "code": "expired"}) raise HTTPException(410, detail={"detail": "Share link expired", "code": "expired"})
# Check password # Check password
if link.password_hash is not None: if link.password_hash is not None:
# Password must be provided via query param or header — check query param # Password must be provided via query param or header — check query param
from fastapi import Request
# We need the request to get the password param # We need the request to get the password param
# For simplicity, return that password is required # For simplicity, return that password is required
raise HTTPException( raise HTTPException(
@@ -267,25 +276,27 @@ async def public_access_with_password(
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
): ):
"""Public share access with password verification — no auth required.""" """Public share access with password verification — no auth required."""
result = await db.execute( result = await db.execute(select(ShareLink).where(ShareLink.token == token))
select(ShareLink).where(ShareLink.token == token)
)
link = result.scalar_one_or_none() link = result.scalar_one_or_none()
if link is None: if link is None:
raise HTTPException(404, detail={"detail": "Share link not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Share link not found", "code": "not_found"})
# Check expiry # Check expiry
if link.expires_at is not None: if link.expires_at is not None:
now = datetime.now(timezone.utc) now = datetime.now(UTC)
if now > link.expires_at: if now > link.expires_at:
raise HTTPException(410, detail={"detail": "Share link expired", "code": "expired"}) raise HTTPException(410, detail={"detail": "Share link expired", "code": "expired"})
# Verify password if required # Verify password if required
if link.password_hash is not None: if link.password_hash is not None:
if body.password is None: if body.password is None:
raise HTTPException(401, detail={"detail": "Password required", "code": "password_required"}) raise HTTPException(
401, detail={"detail": "Password required", "code": "password_required"}
)
if not verify_password(body.password, link.password_hash): if not verify_password(body.password, link.password_hash):
raise HTTPException(403, detail={"detail": "Invalid password", "code": "invalid_password"}) raise HTTPException(
403, detail={"detail": "Invalid password", "code": "invalid_password"}
)
return { return {
"file_id": str(link.file_id), "file_id": str(link.file_id),
+5 -2
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import uuid import uuid
from sqlalchemy import String, ForeignKey, Index, UniqueConstraint from sqlalchemy import ForeignKey, Index, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID as PGUUID from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
@@ -33,7 +33,10 @@ class TagAssignment(Base, TenantMixin):
__tablename__ = "tag_assignments" __tablename__ = "tag_assignments"
__table_args__ = ( __table_args__ = (
UniqueConstraint( UniqueConstraint(
"tenant_id", "tag_id", "entity_type", "entity_id", "tenant_id",
"tag_id",
"entity_type",
"entity_id",
name="uq_tag_assignments_entity", name="uq_tag_assignments_entity",
), ),
Index("ix_tag_assignments_tag", "tenant_id", "tag_id"), Index("ix_tag_assignments_tag", "tenant_id", "tag_id"),
+41 -28
View File
@@ -5,18 +5,18 @@ from __future__ import annotations
import uuid import uuid
from fastapi import APIRouter, Body, Depends, HTTPException, Response, status from fastapi import APIRouter, Body, Depends, HTTPException, Response, status
from sqlalchemy import select, func, delete from sqlalchemy import delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.db import get_db from app.core.db import get_db
from app.deps import get_current_user from app.deps import get_current_user
from app.plugins.builtins.tags.models import Tag, TagAssignment from app.plugins.builtins.tags.models import Tag, TagAssignment
from app.plugins.builtins.tags.schemas import ( from app.plugins.builtins.tags.schemas import (
TagCreate,
TagUpdate,
TagAssignRequest, TagAssignRequest,
TagUnassignRequest,
TagBulkAssignRequest, TagBulkAssignRequest,
TagCreate,
TagUnassignRequest,
TagUpdate,
) )
router = APIRouter(prefix="/api/v1/tags", tags=["tags"]) router = APIRouter(prefix="/api/v1/tags", tags=["tags"])
@@ -28,7 +28,9 @@ def _parse_uuid(val: str, field: str) -> uuid.UUID:
try: try:
return uuid.UUID(val) return uuid.UUID(val)
except (ValueError, TypeError): except (ValueError, TypeError):
raise HTTPException(400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": f"Invalid {field}", "code": "invalid_id"}
) from None
@router.get("") @router.get("")
@@ -107,9 +109,7 @@ async def update_tag(
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
tid = _parse_uuid(tag_id, "tag_id") tid = _parse_uuid(tag_id, "tag_id")
result = await db.execute( result = await db.execute(select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id))
select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id)
)
tag = result.scalar_one_or_none() tag = result.scalar_one_or_none()
if tag is None: if tag is None:
raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"})
@@ -122,7 +122,9 @@ async def update_tag(
select(Tag).where(Tag.tenant_id == tenant_id, Tag.name == data["name"]) select(Tag).where(Tag.tenant_id == tenant_id, Tag.name == data["name"])
) )
if dup.scalar_one_or_none() is not None: if dup.scalar_one_or_none() is not None:
raise HTTPException(409, detail={"detail": "Tag name already exists", "code": "duplicate"}) raise HTTPException(
409, detail={"detail": "Tag name already exists", "code": "duplicate"}
)
tag.name = data["name"] tag.name = data["name"]
if "color" in data: if "color" in data:
tag.color = data["color"] tag.color = data["color"]
@@ -148,12 +150,12 @@ async def assign_tag(
entity_id = _parse_uuid(body.entity_id, "entity_id") entity_id = _parse_uuid(body.entity_id, "entity_id")
if body.entity_type not in VALID_ENTITY_TYPES: if body.entity_type not in VALID_ENTITY_TYPES:
raise HTTPException(400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}) raise HTTPException(
400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}
)
# Verify tag exists # Verify tag exists
tag_result = await db.execute( tag_result = await db.execute(select(Tag).where(Tag.id == tag_id, Tag.tenant_id == tenant_id))
select(Tag).where(Tag.id == tag_id, Tag.tenant_id == tenant_id)
)
if tag_result.scalar_one_or_none() is None: if tag_result.scalar_one_or_none() is None:
raise HTTPException(404, detail={"detail": "Tag not found", "code": "tag_not_found"}) raise HTTPException(404, detail={"detail": "Tag not found", "code": "tag_not_found"})
@@ -167,7 +169,13 @@ async def assign_tag(
) )
) )
if existing.scalar_one_or_none() is not None: if existing.scalar_one_or_none() is not None:
return {"id": str(tag_id), "tag_id": str(tag_id), "entity_type": body.entity_type, "entity_id": str(entity_id), "already_assigned": True} return {
"id": str(tag_id),
"tag_id": str(tag_id),
"entity_type": body.entity_type,
"entity_id": str(entity_id),
"already_assigned": True,
}
assignment = TagAssignment( assignment = TagAssignment(
tenant_id=tenant_id, tenant_id=tenant_id,
@@ -177,7 +185,13 @@ async def assign_tag(
) )
db.add(assignment) db.add(assignment)
await db.flush() await db.flush()
return {"id": str(assignment.id), "tag_id": str(tag_id), "entity_type": body.entity_type, "entity_id": str(entity_id), "already_assigned": False} return {
"id": str(assignment.id),
"tag_id": str(tag_id),
"entity_type": body.entity_type,
"entity_id": str(entity_id),
"already_assigned": False,
}
@router.delete("/assign", status_code=status.HTTP_204_NO_CONTENT) @router.delete("/assign", status_code=status.HTTP_204_NO_CONTENT)
@@ -217,22 +231,21 @@ async def delete_tag(
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
tid = _parse_uuid(tag_id, "tag_id") tid = _parse_uuid(tag_id, "tag_id")
result = await db.execute( result = await db.execute(select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id))
select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id)
)
tag = result.scalar_one_or_none() tag = result.scalar_one_or_none()
if tag is None: if tag is None:
raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"})
# Cascade delete assignments # Cascade delete assignments
await db.execute( await db.execute(
delete(TagAssignment).where(TagAssignment.tag_id == tid, TagAssignment.tenant_id == tenant_id) delete(TagAssignment).where(
TagAssignment.tag_id == tid, TagAssignment.tenant_id == tenant_id
)
) )
await db.delete(tag) await db.delete(tag)
return Response(status_code=status.HTTP_204_NO_CONTENT) return Response(status_code=status.HTTP_204_NO_CONTENT)
@router.post("/bulk-assign") @router.post("/bulk-assign")
async def bulk_assign_tags( async def bulk_assign_tags(
body: TagBulkAssignRequest, body: TagBulkAssignRequest,
@@ -244,17 +257,19 @@ async def bulk_assign_tags(
entity_id = _parse_uuid(body.entity_id, "entity_id") entity_id = _parse_uuid(body.entity_id, "entity_id")
if body.entity_type not in VALID_ENTITY_TYPES: if body.entity_type not in VALID_ENTITY_TYPES:
raise HTTPException(400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}) raise HTTPException(
400, detail={"detail": "Invalid entity_type", "code": "invalid_entity_type"}
)
tag_ids = [_parse_uuid(tid, "tag_id") for tid in body.tag_ids] tag_ids = [_parse_uuid(tid, "tag_id") for tid in body.tag_ids]
# Verify all tags exist # Verify all tags exist
for tid in tag_ids: for tid in tag_ids:
result = await db.execute( result = await db.execute(select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id))
select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id)
)
if result.scalar_one_or_none() is None: if result.scalar_one_or_none() is None:
raise HTTPException(404, detail={"detail": f"Tag {tid} not found", "code": "tag_not_found"}) raise HTTPException(
404, detail={"detail": f"Tag {tid} not found", "code": "tag_not_found"}
)
assigned = [] assigned = []
already = [] already = []
@@ -299,9 +314,7 @@ async def list_tag_entities(
tid = _parse_uuid(tag_id, "tag_id") tid = _parse_uuid(tag_id, "tag_id")
# Verify tag exists # Verify tag exists
tag_result = await db.execute( tag_result = await db.execute(select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id))
select(Tag).where(Tag.id == tid, Tag.tenant_id == tenant_id)
)
if tag_result.scalar_one_or_none() is None: if tag_result.scalar_one_or_none() is None:
raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Tag not found", "code": "not_found"})
-2
View File
@@ -2,8 +2,6 @@
from __future__ import annotations from __future__ import annotations
import uuid
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
+61 -15
View File
@@ -10,21 +10,35 @@ class PluginRouteDef(BaseModel):
path: str = Field(..., description="URL path prefix, e.g. /api/v1/plugin-mail") path: str = Field(..., description="URL path prefix, e.g. /api/v1/plugin-mail")
module: str = Field(..., description="Dotted path to the module containing the APIRouter") module: str = Field(..., description="Dotted path to the module containing the APIRouter")
router_attr: str = Field(default="router", description="Attribute name of the APIRouter in the module") router_attr: str = Field(
default="router", description="Attribute name of the APIRouter in the module"
)
class PluginManifest(BaseModel): class PluginManifest(BaseModel):
"""Manifest describing a plugin's metadata, dependencies, and capabilities.""" """Manifest describing a plugin's metadata, dependencies, and capabilities."""
name: str = Field(..., min_length=1, max_length=80, description="Unique plugin identifier (snake_case)") name: str = Field(
..., min_length=1, max_length=80, description="Unique plugin identifier (snake_case)"
)
version: str = Field(..., min_length=1, max_length=40, description="Semantic version string") version: str = Field(..., min_length=1, max_length=40, description="Semantic version string")
display_name: str = Field(..., min_length=1, max_length=120) display_name: str = Field(..., min_length=1, max_length=120)
description: str = Field(default="", max_length=500) description: str = Field(default="", max_length=500)
dependencies: list[str] = Field(default_factory=list, description="Other plugin names this plugin depends on") dependencies: list[str] = Field(
routes: list[PluginRouteDef] = Field(default_factory=list, description="Route definitions to register on activation") default_factory=list, description="Other plugin names this plugin depends on"
events: list[str] = Field(default_factory=list, description="Event names this plugin listens to") )
migrations: list[str] = Field(default_factory=list, description="Migration file names (ordered, e.g. 0001_initial.sql)") routes: list[PluginRouteDef] = Field(
permissions: list[str] = Field(default_factory=list, description="Required permissions for this plugin") default_factory=list, description="Route definitions to register on activation"
)
events: list[str] = Field(
default_factory=list, description="Event names this plugin listens to"
)
migrations: list[str] = Field(
default_factory=list, description="Migration file names (ordered, e.g. 0001_initial.sql)"
)
permissions: list[str] = Field(
default_factory=list, description="Required permissions for this plugin"
)
@field_validator("name") @field_validator("name")
@classmethod @classmethod
@@ -46,15 +60,47 @@ class ManifestSchemaResponse(BaseModel):
# Pre-built schema documentation for GET /api/v1/plugins/manifest endpoint # Pre-built schema documentation for GET /api/v1/plugins/manifest endpoint
MANIFEST_SCHEMA_DOC = ManifestSchemaResponse( MANIFEST_SCHEMA_DOC = ManifestSchemaResponse(
fields={ fields={
"name": {"type": "str", "required": "true", "description": "Unique plugin identifier (snake_case, max 80 chars)"}, "name": {
"type": "str",
"required": "true",
"description": "Unique plugin identifier (snake_case, max 80 chars)",
},
"version": {"type": "str", "required": "true", "description": "Semantic version string"}, "version": {"type": "str", "required": "true", "description": "Semantic version string"},
"display_name": {"type": "str", "required": "true", "description": "Human-readable plugin name"}, "display_name": {
"description": {"type": "str", "required": "false", "description": "Plugin description (max 500 chars)"}, "type": "str",
"dependencies": {"type": "list[str]", "required": "false", "description": "Other plugin names required"}, "required": "true",
"routes": {"type": "list[PluginRouteDef]", "required": "false", "description": "Route definitions to register"}, "description": "Human-readable plugin name",
"events": {"type": "list[str]", "required": "false", "description": "Event names to listen to"}, },
"migrations": {"type": "list[str]", "required": "false", "description": "Migration file names (ordered)"}, "description": {
"permissions": {"type": "list[str]", "required": "false", "description": "Required permissions"}, "type": "str",
"required": "false",
"description": "Plugin description (max 500 chars)",
},
"dependencies": {
"type": "list[str]",
"required": "false",
"description": "Other plugin names required",
},
"routes": {
"type": "list[PluginRouteDef]",
"required": "false",
"description": "Route definitions to register",
},
"events": {
"type": "list[str]",
"required": "false",
"description": "Event names to listen to",
},
"migrations": {
"type": "list[str]",
"required": "false",
"description": "Migration file names (ordered)",
},
"permissions": {
"type": "list[str]",
"required": "false",
"description": "Required permissions",
},
}, },
example=PluginManifest( example=PluginManifest(
name="example_plugin", name="example_plugin",
+7 -2
View File
@@ -6,7 +6,7 @@ import os
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
from sqlalchemy import text, inspect from sqlalchemy import inspect, text
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
from app.models.plugin import PluginMigration from app.models.plugin import PluginMigration
@@ -21,7 +21,9 @@ class MigrationRunner:
def __init__(self, engine: AsyncEngine, migrations_dir: str | None = None) -> None: def __init__(self, engine: AsyncEngine, migrations_dir: str | None = None) -> None:
self._engine = engine self._engine = engine
self._migrations_dir = Path(migrations_dir or os.path.join(os.path.dirname(__file__), "migrations")) self._migrations_dir = Path(
migrations_dir or os.path.join(os.path.dirname(__file__), "migrations")
)
def _resolve_migration_path(self, filename: str, plugin_name: str | None = None) -> Path: def _resolve_migration_path(self, filename: str, plugin_name: str | None = None) -> Path:
"""Resolve a migration filename to an absolute path. """Resolve a migration filename to an absolute path.
@@ -213,6 +215,7 @@ class MigrationRunner:
async def _get_table_names(self) -> set[str]: async def _get_table_names(self) -> set[str]:
"""Get current table names from the database (separate connection).""" """Get current table names from the database (separate connection)."""
def _get_names(sync_conn): def _get_names(sync_conn):
insp = inspect(sync_conn) insp = inspect(sync_conn)
return set(insp.get_table_names()) return set(insp.get_table_names())
@@ -222,6 +225,7 @@ class MigrationRunner:
async def _table_has_column(self, table_name: str, column_name: str) -> bool: async def _table_has_column(self, table_name: str, column_name: str) -> bool:
"""Check if a table has a specific column (separate connection).""" """Check if a table has a specific column (separate connection)."""
def _has_col(sync_conn): def _has_col(sync_conn):
insp = inspect(sync_conn) insp = inspect(sync_conn)
if table_name not in insp.get_table_names(): if table_name not in insp.get_table_names():
@@ -289,6 +293,7 @@ class MigrationRunner:
def _extract_table_names(sql: str) -> list[str]: def _extract_table_names(sql: str) -> list[str]:
"""Extract table names from CREATE TABLE statements in SQL.""" """Extract table names from CREATE TABLE statements in SQL."""
import re import re
# Match CREATE TABLE [IF NOT EXISTS] "table_name" or CREATE TABLE table_name # Match CREATE TABLE [IF NOT EXISTS] "table_name" or CREATE TABLE table_name
pattern = r'CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?["\']?(\w+)["\']?' pattern = r'CREATE\s+TABLE\s+(?:IF\s+NOT\s+EXISTS\s+)?["\']?(\w+)["\']?'
return re.findall(pattern, sql, re.IGNORECASE) return re.findall(pattern, sql, re.IGNORECASE)
+52 -48
View File
@@ -8,7 +8,7 @@ from typing import Any
from fastapi import FastAPI from fastapi import FastAPI
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession, AsyncEngine from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
from app.core.event_bus import EventBus, get_event_bus from app.core.event_bus import EventBus, get_event_bus
from app.core.service_container import ServiceContainer, get_container from app.core.service_container import ServiceContainer, get_container
@@ -78,7 +78,8 @@ class PluginRegistry:
return discovered return discovered
import pkgutil import pkgutil
for importer, modname, ispkg in pkgutil.iter_modules(pkg_path):
for _importer, modname, _ispkg in pkgutil.iter_modules(pkg_path):
if modname.startswith("_"): if modname.startswith("_"):
continue continue
try: try:
@@ -144,9 +145,7 @@ class PluginRegistry:
# Run migrations # Run migrations
if plugin.manifest.migrations: if plugin.manifest.migrations:
await self.migration_runner.run_all_migrations( await self.migration_runner.run_all_migrations(db, name, plugin.manifest.migrations)
db, name, plugin.manifest.migrations
)
# Call on_install hook # Call on_install hook
await plugin.on_install(db, self._container) await plugin.on_install(db, self._container)
@@ -234,7 +233,8 @@ class PluginRegistry:
paths_to_remove.add(route.path) paths_to_remove.add(route.path)
# Remove matching routes from app # Remove matching routes from app
self._app.router.routes = [ self._app.router.routes = [
r for r in self._app.router.routes r
for r in self._app.router.routes
if not (hasattr(r, "path") and r.path in paths_to_remove) if not (hasattr(r, "path") and r.path in paths_to_remove)
] ]
@@ -309,50 +309,56 @@ class PluginRegistry:
for name, plugin in self._plugins.items(): for name, plugin in self._plugins.items():
record = db_records.get(name) record = db_records.get(name)
if record is not None: if record is not None:
plugins_list.append({ plugins_list.append(
"name": name, {
"display_name": record.display_name, "name": name,
"version": record.version, "display_name": record.display_name,
"status": record.status, "version": record.version,
"installed": record.installed, "status": record.status,
"active": record.active, "installed": record.installed,
"description": plugin.manifest.description, "active": record.active,
"dependencies": plugin.manifest.dependencies, "description": plugin.manifest.description,
"events": plugin.manifest.events, "dependencies": plugin.manifest.dependencies,
"migrations": plugin.manifest.migrations, "events": plugin.manifest.events,
"permissions": plugin.manifest.permissions, "migrations": plugin.manifest.migrations,
}) "permissions": plugin.manifest.permissions,
}
)
else: else:
plugins_list.append({ plugins_list.append(
"name": name, {
"display_name": plugin.manifest.display_name, "name": name,
"version": plugin.version, "display_name": plugin.manifest.display_name,
"status": "discovered", "version": plugin.version,
"installed": False, "status": "discovered",
"active": False, "installed": False,
"description": plugin.manifest.description, "active": False,
"dependencies": plugin.manifest.dependencies, "description": plugin.manifest.description,
"events": plugin.manifest.events, "dependencies": plugin.manifest.dependencies,
"migrations": plugin.manifest.migrations, "events": plugin.manifest.events,
"permissions": plugin.manifest.permissions, "migrations": plugin.manifest.migrations,
}) "permissions": plugin.manifest.permissions,
}
)
# Also include DB-only records (plugins that were installed but no longer discovered) # Also include DB-only records (plugins that were installed but no longer discovered)
for name, record in db_records.items(): for name, record in db_records.items():
if name not in self._plugins: if name not in self._plugins:
plugins_list.append({ plugins_list.append(
"name": name, {
"display_name": record.display_name, "name": name,
"version": record.version, "display_name": record.display_name,
"status": record.status, "version": record.version,
"installed": record.installed, "status": record.status,
"active": record.active, "installed": record.installed,
"description": "", "active": record.active,
"dependencies": [], "description": "",
"events": [], "dependencies": [],
"migrations": [], "events": [],
"permissions": [], "migrations": [],
}) "permissions": [],
}
)
return plugins_list return plugins_list
@@ -360,9 +366,7 @@ class PluginRegistry:
async def _get_plugin_record(self, db: AsyncSession, name: str) -> PluginModel | None: async def _get_plugin_record(self, db: AsyncSession, name: str) -> PluginModel | None:
"""Fetch a plugin record from DB by name.""" """Fetch a plugin record from DB by name."""
result = await db.execute( result = await db.execute(select(PluginModel).where(PluginModel.name == name))
select(PluginModel).where(PluginModel.name == name)
)
return result.scalar_one_or_none() return result.scalar_one_or_none()
+14 -1
View File
@@ -1,3 +1,16 @@
"""Routes package.""" """Routes package."""
from app.routes import auth, users, roles, tenants, health, notifications, companies, contacts, import_export, plugins, ai_copilot, workflows from app.routes import (
ai_copilot, # noqa: F401
auth, # noqa: F401
companies, # noqa: F401
contacts, # noqa: F401
health, # noqa: F401
import_export, # noqa: F401
notifications, # noqa: F401
plugins, # noqa: F401
roles, # noqa: F401
tenants, # noqa: F401
users, # noqa: F401
workflows, # noqa: F401
)
+13 -6
View File
@@ -7,12 +7,11 @@ import uuid
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query, status
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.auth import check_permission
from app.core.db import get_db from app.core.db import get_db
from app.deps import get_current_user from app.deps import get_current_user
from app.schemas.ai_copilot import ( from app.schemas.ai_copilot import (
CopilotQueryRequest,
CopilotExecuteRequest, CopilotExecuteRequest,
CopilotQueryRequest,
) )
from app.services import ai_copilot_service from app.services import ai_copilot_service
@@ -34,7 +33,9 @@ async def copilot_query(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
result = await ai_copilot_service.process_query( result = await ai_copilot_service.process_query(
db, tenant_id, user_id, db,
tenant_id,
user_id,
query=body.query, query=body.query,
conversation_id=body.conversation_id, conversation_id=body.conversation_id,
context=body.context, context=body.context,
@@ -65,7 +66,10 @@ async def copilot_execute(
role = current_user.get("role", "viewer") role = current_user.get("role", "viewer")
result = await ai_copilot_service.execute_action( result = await ai_copilot_service.execute_action(
db, tenant_id, user_id, role, db,
tenant_id,
user_id,
role,
conversation_id=body.conversation_id, conversation_id=body.conversation_id,
action=body.action.model_dump(), action=body.action.model_dump(),
) )
@@ -97,6 +101,9 @@ async def copilot_history(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
return await ai_copilot_service.get_history( return await ai_copilot_service.get_history(
db, tenant_id, user_id, db,
page=page, page_size=page_size, tenant_id,
user_id,
page=page,
page_size=page_size,
) )
+10 -8
View File
@@ -3,19 +3,19 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings from app.config import get_settings
from app.core.rate_limit import check_rate_limit, get_client_ip, reset_rate_limit
from app.core.auth import get_redis from app.core.auth import get_redis
from app.core.db import get_db from app.core.db import get_db
from app.deps import get_current_user from app.core.rate_limit import check_rate_limit, get_client_ip, reset_rate_limit
from app.schemas.auth import ( from app.schemas.auth import (
LoginRequest, PasswordResetConfirm, PasswordResetRequest, LoginRequest,
SwitchTenantRequest, AuthResponse, PasswordResetConfirm,
PasswordResetRequest,
SwitchTenantRequest,
) )
from app.services.auth_service import auth_service from app.services.auth_service import auth_service
@@ -64,6 +64,7 @@ async def login(
) )
# Return auth info in body # Return auth info in body
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
resp = JSONResponse( resp = JSONResponse(
status_code=status.HTTP_200_OK, status_code=status.HTTP_200_OK,
content={ content={
@@ -100,6 +101,7 @@ async def logout(
await auth_service.logout(redis, session_id) await auth_service.logout(redis, session_id)
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
resp = JSONResponse( resp = JSONResponse(
status_code=status.HTTP_200_OK, status_code=status.HTTP_200_OK,
content={"message": "Logged out"}, content={"message": "Logged out"},
@@ -153,7 +155,7 @@ async def switch_tenant(
raise HTTPException( raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, status_code=status.HTTP_400_BAD_REQUEST,
detail={"detail": "Invalid tenant_id", "code": "invalid_tenant_id"}, detail={"detail": "Invalid tenant_id", "code": "invalid_tenant_id"},
) ) from None
result = await auth_service.switch_tenant(db, redis, session_id, new_tenant_id) result = await auth_service.switch_tenant(db, redis, session_id, new_tenant_id)
if result is None: if result is None:
@@ -180,7 +182,7 @@ async def password_reset_request(
): ):
"""Request password reset. Always returns 200 (no user enumeration).""" """Request password reset. Always returns 200 (no user enumeration)."""
ip = get_client_ip(request) ip = get_client_ip(request)
redis = get_redis() redis = get_redis() # noqa: F841
await check_rate_limit( await check_rate_limit(
f"auth:reset:{ip}", f"auth:reset:{ip}",
@@ -200,7 +202,7 @@ async def password_reset_confirm(
): ):
"""Reset password with a valid token.""" """Reset password with a valid token."""
ip = get_client_ip(request) ip = get_client_ip(request)
redis = get_redis() redis = get_redis() # noqa: F841
await check_rate_limit( await check_rate_limit(
f"auth:reset_confirm:{ip}", f"auth:reset_confirm:{ip}",
+38 -14
View File
@@ -30,10 +30,14 @@ async def list_companies(
"""List companies with pagination, FTS search, industry filter, and sorting.""" """List companies with pagination, FTS search, industry filter, and sorting."""
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
result = await company_service.list_companies( result = await company_service.list_companies(
db, tenant_id, db,
page=page, page_size=page_size, tenant_id,
search=search, industry=industry, page=page,
sort_by=sort_by, sort_order=sort_order, page_size=page_size,
search=search,
industry=industry,
sort_by=sort_by,
sort_order=sort_order,
) )
return result return result
@@ -72,7 +76,10 @@ async def export_companies(
if format == "csv": if format == "csv":
csv_data = await company_service.export_companies_csv( csv_data = await company_service.export_companies_csv(
db, tenant_id, industry=industry, search=search, db,
tenant_id,
industry=industry,
search=search,
) )
return Response( return Response(
content=csv_data, content=csv_data,
@@ -81,7 +88,10 @@ async def export_companies(
) )
elif format == "xlsx": elif format == "xlsx":
xlsx_data = await company_service.export_companies_xlsx( xlsx_data = await company_service.export_companies_xlsx(
db, tenant_id, industry=industry, search=search, db,
tenant_id,
industry=industry,
search=search,
) )
return Response( return Response(
content=xlsx_data, content=xlsx_data,
@@ -102,7 +112,9 @@ async def get_company(
try: try:
cid = uuid.UUID(company_id) cid = uuid.UUID(company_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid company_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid company_id", "code": "invalid_id"}
) from None
data = await company_service.get_company_detail(db, tenant_id, cid) data = await company_service.get_company_detail(db, tenant_id, cid)
if data is None: if data is None:
@@ -128,7 +140,9 @@ async def update_company(
try: try:
cid = uuid.UUID(company_id) cid = uuid.UUID(company_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid company_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid company_id", "code": "invalid_id"}
) from None
data = body.model_dump(exclude_unset=True) data = body.model_dump(exclude_unset=True)
result = await company_service.update_company(db, tenant_id, user_id, cid, data) result = await company_service.update_company(db, tenant_id, user_id, cid, data)
@@ -155,9 +169,13 @@ async def delete_company(
try: try:
cid = uuid.UUID(company_id) cid = uuid.UUID(company_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid company_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid company_id", "code": "invalid_id"}
) from None
deleted = await company_service.soft_delete_company(db, tenant_id, user_id, cid, cascade=cascade) deleted = await company_service.soft_delete_company(
db, tenant_id, user_id, cid, cascade=cascade
)
if not deleted: if not deleted:
raise HTTPException(404, detail={"detail": "Company not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Company not found", "code": "not_found"})
@@ -183,11 +201,13 @@ async def link_company_contact(
comp_id = uuid.UUID(company_id) comp_id = uuid.UUID(company_id)
cont_id = uuid.UUID(contact_id) cont_id = uuid.UUID(contact_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) from None
result = await company_service.link_contact(db, tenant_id, user_id, comp_id, cont_id) result = await company_service.link_contact(db, tenant_id, user_id, comp_id, cont_id)
if result is None: if result is None:
raise HTTPException(404, detail={"detail": "Company or contact not found", "code": "not_found"}) raise HTTPException(
404, detail={"detail": "Company or contact not found", "code": "not_found"}
)
return result return result
@@ -210,7 +230,7 @@ async def unlink_company_contact(
comp_id = uuid.UUID(company_id) comp_id = uuid.UUID(company_id)
cont_id = uuid.UUID(contact_id) cont_id = uuid.UUID(contact_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) from None
unlinked = await company_service.unlink_contact(db, tenant_id, user_id, comp_id, cont_id) unlinked = await company_service.unlink_contact(db, tenant_id, user_id, comp_id, cont_id)
if not unlinked: if not unlinked:
@@ -230,11 +250,15 @@ async def get_company_emails(
try: try:
cid = uuid.UUID(company_id) cid = uuid.UUID(company_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid company_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid company_id", "code": "invalid_id"}
) from None
# Verify company exists # Verify company exists
from sqlalchemy import select from sqlalchemy import select
from app.models.company import Company from app.models.company import Company
q = select(Company).where( q = select(Company).where(
Company.id == cid, Company.id == cid,
Company.tenant_id == tenant_id, Company.tenant_id == tenant_id,
+16 -6
View File
@@ -29,9 +29,13 @@ async def list_contacts(
"""List contacts with pagination and optional search.""" """List contacts with pagination and optional search."""
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
result = await contact_service.list_contacts( result = await contact_service.list_contacts(
db, tenant_id, db,
page=page, page_size=page_size, tenant_id,
search=search, sort_by=sort_by, sort_order=sort_order, page=page,
page_size=page_size,
search=search,
sort_by=sort_by,
sort_order=sort_order,
) )
return result return result
@@ -68,7 +72,9 @@ async def get_contact(
try: try:
cid = uuid.UUID(contact_id) cid = uuid.UUID(contact_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}
) from None
data = await contact_service.get_contact_detail(db, tenant_id, cid) data = await contact_service.get_contact_detail(db, tenant_id, cid)
if data is None: if data is None:
@@ -94,7 +100,9 @@ async def update_contact(
try: try:
cid = uuid.UUID(contact_id) cid = uuid.UUID(contact_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}
) from None
data = body.model_dump(exclude_unset=True) data = body.model_dump(exclude_unset=True)
result = await contact_service.update_contact(db, tenant_id, user_id, cid, data) result = await contact_service.update_contact(db, tenant_id, user_id, cid, data)
@@ -121,7 +129,9 @@ async def delete_contact(
try: try:
cid = uuid.UUID(contact_id) cid = uuid.UUID(contact_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid contact_id", "code": "invalid_id"}
) from None
if gdpr: if gdpr:
deleted = await contact_service.gdpr_hard_delete_contact(db, tenant_id, user_id, cid) deleted = await contact_service.gdpr_hard_delete_contact(db, tenant_id, user_id, cid)
+20 -5
View File
@@ -3,16 +3,21 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from typing import Any
from fastapi import APIRouter, Depends, HTTPException, UploadFile, File, Form, Query, Response, status from fastapi import (
APIRouter,
Depends,
File,
Form,
HTTPException,
UploadFile,
)
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.auth import check_permission from app.core.auth import check_permission
from app.core.db import get_db from app.core.db import get_db
from app.deps import get_current_user from app.deps import get_current_user
from app.services import import_export_service from app.services import import_export_service
from app.services import company_service
router = APIRouter(prefix="/api/v1", tags=["import_export"]) router = APIRouter(prefix="/api/v1", tags=["import_export"])
@@ -38,7 +43,12 @@ async def import_csv(
csv_content = content.decode("utf-8") csv_content = content.decode("utf-8")
result = await import_export_service.import_csv( result = await import_export_service.import_csv(
db, tenant_id, user_id, csv_content, entity_type=entity_type, dry_run=False, db,
tenant_id,
user_id,
csv_content,
entity_type=entity_type,
dry_run=False,
) )
return result return result
@@ -62,6 +72,11 @@ async def import_csv_preview(
csv_content = content.decode("utf-8") csv_content = content.decode("utf-8")
result = await import_export_service.import_csv( result = await import_export_service.import_csv(
db, tenant_id, user_id, csv_content, entity_type=entity_type, dry_run=True, db,
tenant_id,
user_id,
csv_content,
entity_type=entity_type,
dry_run=True,
) )
return result return result
+7 -3
View File
@@ -4,12 +4,14 @@ from __future__ import annotations
import uuid import uuid
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.db import get_db from app.core.db import get_db
from app.core.notifications import ( from app.core.notifications import (
get_unread_count, list_notifications, mark_notification_read, get_unread_count,
list_notifications,
mark_notification_read,
) )
from app.deps import get_current_user from app.deps import get_current_user
@@ -41,7 +43,9 @@ async def mark_notification_read_endpoint(
try: try:
nid = uuid.UUID(notification_id) nid = uuid.UUID(notification_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid notification_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid notification_id", "code": "invalid_id"}
) from None
notif = await mark_notification_read(db, tenant_id, user_id, nid) notif = await mark_notification_read(db, tenant_id, user_id, nid)
if notif is None: if notif is None:
+35 -15
View File
@@ -2,7 +2,7 @@
from __future__ import annotations from __future__ import annotations
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.db import get_db from app.core.db import get_db
@@ -44,20 +44,26 @@ async def install_plugin(
Idempotent: returns 200 if already installed. Idempotent: returns 200 if already installed.
""" """
import uuid as uuid_mod import uuid as uuid_mod
service = get_plugin_service() service = get_plugin_service()
try: try:
result = await service.install_plugin( result = await service.install_plugin(
db, name, db,
name,
tenant_id=uuid_mod.UUID(current_user["tenant_id"]), tenant_id=uuid_mod.UUID(current_user["tenant_id"]),
user_id=uuid_mod.UUID(current_user["user_id"]), user_id=uuid_mod.UUID(current_user["user_id"]),
) )
return result return result
except ValueError as exc: except ValueError as exc:
if "not found" in str(exc).lower(): if "not found" in str(exc).lower():
raise HTTPException(404, detail={"detail": str(exc), "code": "plugin_not_found"}) raise HTTPException(
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) 404, detail={"detail": str(exc), "code": "plugin_not_found"}
) from None
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) from None
except MigrationValidationError as exc: except MigrationValidationError as exc:
raise HTTPException(422, detail={"detail": str(exc), "code": "migration_validation_error"}) raise HTTPException(
422, detail={"detail": str(exc), "code": "migration_validation_error"}
) from None
@router.post("/{name}/activate") @router.post("/{name}/activate")
@@ -71,20 +77,26 @@ async def activate_plugin(
Idempotent: returns 200 if already active. Idempotent: returns 200 if already active.
""" """
import uuid as uuid_mod import uuid as uuid_mod
service = get_plugin_service() service = get_plugin_service()
try: try:
result = await service.activate_plugin( result = await service.activate_plugin(
db, name, db,
name,
tenant_id=uuid_mod.UUID(current_user["tenant_id"]), tenant_id=uuid_mod.UUID(current_user["tenant_id"]),
user_id=uuid_mod.UUID(current_user["user_id"]), user_id=uuid_mod.UUID(current_user["user_id"]),
) )
return result return result
except ValueError as exc: except ValueError as exc:
if "not installed" in str(exc).lower(): if "not installed" in str(exc).lower():
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_not_installed"}) raise HTTPException(
400, detail={"detail": str(exc), "code": "plugin_not_installed"}
) from None
if "not found" in str(exc).lower(): if "not found" in str(exc).lower():
raise HTTPException(404, detail={"detail": str(exc), "code": "plugin_not_found"}) raise HTTPException(
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) 404, detail={"detail": str(exc), "code": "plugin_not_found"}
) from None
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) from None
@router.post("/{name}/deactivate") @router.post("/{name}/deactivate")
@@ -98,18 +110,22 @@ async def deactivate_plugin(
Idempotent: returns 200 if already inactive. Idempotent: returns 200 if already inactive.
""" """
import uuid as uuid_mod import uuid as uuid_mod
service = get_plugin_service() service = get_plugin_service()
try: try:
result = await service.deactivate_plugin( result = await service.deactivate_plugin(
db, name, db,
name,
tenant_id=uuid_mod.UUID(current_user["tenant_id"]), tenant_id=uuid_mod.UUID(current_user["tenant_id"]),
user_id=uuid_mod.UUID(current_user["user_id"]), user_id=uuid_mod.UUID(current_user["user_id"]),
) )
return result return result
except ValueError as exc: except ValueError as exc:
if "not found" in str(exc).lower(): if "not found" in str(exc).lower():
raise HTTPException(404, detail={"detail": str(exc), "code": "plugin_not_found"}) raise HTTPException(
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) 404, detail={"detail": str(exc), "code": "plugin_not_found"}
) from None
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) from None
@router.delete("/{name}") @router.delete("/{name}")
@@ -124,10 +140,12 @@ async def uninstall_plugin(
Returns 200 with uninstalled status. Idempotent for already-uninstalled plugins. Returns 200 with uninstalled status. Idempotent for already-uninstalled plugins.
""" """
import uuid as uuid_mod import uuid as uuid_mod
service = get_plugin_service() service = get_plugin_service()
try: try:
result = await service.uninstall_plugin( result = await service.uninstall_plugin(
db, name, db,
name,
remove_data=remove_data, remove_data=remove_data,
tenant_id=uuid_mod.UUID(current_user["tenant_id"]), tenant_id=uuid_mod.UUID(current_user["tenant_id"]),
user_id=uuid_mod.UUID(current_user["user_id"]), user_id=uuid_mod.UUID(current_user["user_id"]),
@@ -135,5 +153,7 @@ async def uninstall_plugin(
return result return result
except ValueError as exc: except ValueError as exc:
if "not found" in str(exc).lower(): if "not found" in str(exc).lower():
raise HTTPException(404, detail={"detail": str(exc), "code": "plugin_not_found"}) raise HTTPException(
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) 404, detail={"detail": str(exc), "code": "plugin_not_found"}
) from None
raise HTTPException(400, detail={"detail": str(exc), "code": "plugin_error"}) from None
+18 -6
View File
@@ -4,8 +4,7 @@ from __future__ import annotations
import uuid import uuid
from fastapi import APIRouter, Depends, HTTPException, status from fastapi import APIRouter, Depends, HTTPException, Response, status
from fastapi import Response
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
@@ -37,7 +36,11 @@ async def create_role(
"""Create a custom role (admin only).""" """Create a custom role (admin only)."""
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
role = await role_service.create_role( role = await role_service.create_role(
db, tenant_id, body.name, body.permissions, body.field_permissions, db,
tenant_id,
body.name,
body.permissions,
body.field_permissions,
) )
return JSONResponse( return JSONResponse(
status_code=status.HTTP_201_CREATED, status_code=status.HTTP_201_CREATED,
@@ -62,10 +65,17 @@ async def update_role(
try: try:
rid = uuid.UUID(role_id) rid = uuid.UUID(role_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid role_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid role_id", "code": "invalid_id"}
) from None
role = await role_service.update_role( role = await role_service.update_role(
db, tenant_id, rid, body.name, body.permissions, body.field_permissions, db,
tenant_id,
rid,
body.name,
body.permissions,
body.field_permissions,
) )
if role is None: if role is None:
raise HTTPException(404, detail={"detail": "Role not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "Role not found", "code": "not_found"})
@@ -89,7 +99,9 @@ async def delete_role(
try: try:
rid = uuid.UUID(role_id) rid = uuid.UUID(role_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid role_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid role_id", "code": "invalid_id"}
) from None
success = await role_service.delete_role(db, tenant_id, rid) success = await role_service.delete_role(db, tenant_id, rid)
if not success: if not success:
+6 -4
View File
@@ -4,7 +4,7 @@ from __future__ import annotations
import uuid import uuid
from fastapi import APIRouter, Depends, HTTPException, status from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.db import get_db from app.core.db import get_db
@@ -51,7 +51,9 @@ async def list_tenant_users(
try: try:
tid = uuid.UUID(tenant_id) tid = uuid.UUID(tenant_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid tenant_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid tenant_id", "code": "invalid_id"}
) from None
users = await tenant_service.list_tenant_users(db, tid) users = await tenant_service.list_tenant_users(db, tid)
return {"items": users} return {"items": users}
@@ -69,7 +71,7 @@ async def assign_user_to_tenant(
tid = uuid.UUID(tenant_id) tid = uuid.UUID(tenant_id)
uid = uuid.UUID(body.user_id) uid = uuid.UUID(body.user_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) raise HTTPException(400, detail={"detail": "Invalid ID", "code": "invalid_id"}) from None
ut = await tenant_service.assign_user_to_tenant(db, tid, uid) await tenant_service.assign_user_to_tenant(db, tid, uid)
return {"message": "User assigned to tenant"} return {"message": "User assigned to tenant"}
+40 -11
View File
@@ -5,14 +5,13 @@ from __future__ import annotations
import uuid import uuid
from typing import Any from typing import Any
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from fastapi import Response
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.audit import log_audit from app.core.audit import log_audit
from app.core.db import get_db from app.core.db import get_db
from app.core.notifications import create_notification from app.core.notifications import create_notification
from app.deps import get_current_user, require_admin, get_tenant_id, get_current_user_id from app.deps import get_current_user, require_admin
from app.schemas.user import UserCreate, UserUpdate from app.schemas.user import UserCreate, UserUpdate
from app.services.user_service import user_service from app.services.user_service import user_service
@@ -43,18 +42,32 @@ async def create_user(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
user = await user_service.create_user( user = await user_service.create_user(
db, tenant_id, body.email, body.name, body.password, body.role, body.is_active, db,
tenant_id,
body.email,
body.name,
body.password,
body.role,
body.is_active,
) )
# Audit log # Audit log
await log_audit( await log_audit(
db, tenant_id, user_id, "create", "user", user.id, db,
tenant_id,
user_id,
"create",
"user",
user.id,
changes={"email": body.email, "name": body.name, "role": body.role}, changes={"email": body.email, "name": body.name, "role": body.role},
) )
# Notification # Notification
await create_notification( await create_notification(
db, tenant_id, user.id, "info", db,
tenant_id,
user.id,
"info",
"Account created", "Account created",
f"Your account has been created by {current_user['name']}.", f"Your account has been created by {current_user['name']}.",
) )
@@ -80,7 +93,9 @@ async def get_user(
try: try:
uid = uuid.UUID(user_id) uid = uuid.UUID(user_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid user_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid user_id", "code": "invalid_id"}
) from None
user = await user_service.get_user(db, tenant_id, uid) user = await user_service.get_user(db, tenant_id, uid)
if user is None: if user is None:
@@ -109,7 +124,9 @@ async def update_user(
try: try:
uid = uuid.UUID(user_id) uid = uuid.UUID(user_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid user_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid user_id", "code": "invalid_id"}
) from None
changes: dict[str, Any] = {} changes: dict[str, Any] = {}
if body.name is not None: if body.name is not None:
@@ -120,7 +137,12 @@ async def update_user(
changes["is_active"] = body.is_active changes["is_active"] = body.is_active
user = await user_service.update_user( user = await user_service.update_user(
db, tenant_id, uid, body.name, body.role, body.is_active, db,
tenant_id,
uid,
body.name,
body.role,
body.is_active,
) )
if user is None: if user is None:
raise HTTPException(404, detail={"detail": "User not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "User not found", "code": "not_found"})
@@ -149,7 +171,9 @@ async def delete_user(
try: try:
uid = uuid.UUID(user_id) uid = uuid.UUID(user_id)
except ValueError: except ValueError:
raise HTTPException(400, detail={"detail": "Invalid user_id", "code": "invalid_id"}) raise HTTPException(
400, detail={"detail": "Invalid user_id", "code": "invalid_id"}
) from None
# Get user snapshot for audit before deletion # Get user snapshot for audit before deletion
user = await user_service.get_user(db, tenant_id, uid) user = await user_service.get_user(db, tenant_id, uid)
@@ -161,7 +185,12 @@ async def delete_user(
raise HTTPException(404, detail={"detail": "User not found", "code": "not_found"}) raise HTTPException(404, detail={"detail": "User not found", "code": "not_found"})
await log_audit( await log_audit(
db, tenant_id, acting_user_id, "delete", "user", uid, db,
tenant_id,
acting_user_id,
"delete",
"user",
uid,
changes={"email": user.email, "name": user.name}, changes={"email": user.email, "name": user.name},
) )
+20 -8
View File
@@ -10,7 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.core.auth import check_permission from app.core.auth import check_permission
from app.core.db import get_db from app.core.db import get_db
from app.deps import get_current_user from app.deps import get_current_user
from app.schemas.workflow import WorkflowCreate, WorkflowUpdate, InstanceCreate, AdvanceRequest from app.schemas.workflow import AdvanceRequest, InstanceCreate, WorkflowCreate, WorkflowUpdate
from app.services import workflow_service from app.services import workflow_service
router = APIRouter(prefix="/api/v1/workflows", tags=["workflows"]) router = APIRouter(prefix="/api/v1/workflows", tags=["workflows"])
@@ -18,6 +18,7 @@ router = APIRouter(prefix="/api/v1/workflows", tags=["workflows"])
# ─── Workflow CRUD ─── # ─── Workflow CRUD ───
@router.get("") @router.get("")
async def list_workflows( async def list_workflows(
page: int = Query(1, ge=1), page: int = Query(1, ge=1),
@@ -29,8 +30,10 @@ async def list_workflows(
"""List workflows with pagination.""" """List workflows with pagination."""
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
return await workflow_service.list_workflows( return await workflow_service.list_workflows(
db, tenant_id, db,
page=page, page_size=page_size, tenant_id,
page=page,
page_size=page_size,
is_active=is_active, is_active=is_active,
) )
@@ -67,8 +70,10 @@ async def list_instances(
"""List workflow instances with optional status filter.""" """List workflow instances with optional status filter."""
tenant_id = uuid.UUID(current_user["tenant_id"]) tenant_id = uuid.UUID(current_user["tenant_id"])
return await workflow_service.list_instances( return await workflow_service.list_instances(
db, tenant_id, db,
page=page, page_size=page_size, tenant_id,
page=page,
page_size=page_size,
status_filter=status, status_filter=status,
) )
@@ -146,6 +151,7 @@ async def delete_workflow(
# ─── Instance endpoints ─── # ─── Instance endpoints ───
@router.post("/{workflow_id}/instances", status_code=status.HTTP_201_CREATED) @router.post("/{workflow_id}/instances", status_code=status.HTTP_201_CREATED)
async def create_instance( async def create_instance(
workflow_id: str, workflow_id: str,
@@ -158,7 +164,9 @@ async def create_instance(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
result = await workflow_service.create_instance( result = await workflow_service.create_instance(
db, tenant_id, user_id, db,
tenant_id,
user_id,
workflow_id=workflow_id, workflow_id=workflow_id,
context=body.context, context=body.context,
timeout_hours=body.timeout_hours, timeout_hours=body.timeout_hours,
@@ -204,7 +212,9 @@ async def advance_instance(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
result = await workflow_service.advance_instance( result = await workflow_service.advance_instance(
db, tenant_id, user_id, db,
tenant_id,
user_id,
instance_id=instance_id, instance_id=instance_id,
decision=body.decision, decision=body.decision,
comment=body.comment, comment=body.comment,
@@ -233,7 +243,9 @@ async def cancel_instance(
user_id = uuid.UUID(current_user["user_id"]) user_id = uuid.UUID(current_user["user_id"])
result = await workflow_service.cancel_instance( result = await workflow_service.cancel_instance(
db, tenant_id, user_id, db,
tenant_id,
user_id,
instance_id=instance_id, instance_id=instance_id,
) )
if result is None: if result is None:
+24 -7
View File
@@ -1,13 +1,30 @@
"""Pydantic schemas package.""" """Pydantic schemas package."""
from app.schemas.plugin import PluginInfo, PluginListResponse, PluginActionResponse, PluginUninstallResponse
from app.schemas.ai_copilot import ( from app.schemas.ai_copilot import (
CopilotQueryRequest, CopilotAction, CopilotQueryResponse, CopilotAction, # noqa: F401
CopilotExecuteRequest, CopilotExecuteResponse, CopilotExecuteRequest, # noqa: F401
CopilotHistoryResponse, CopilotMessageResponse, CopilotExecuteResponse, # noqa: F401
CopilotHistoryResponse, # noqa: F401
CopilotMessageResponse, # noqa: F401
CopilotQueryRequest, # noqa: F401
CopilotQueryResponse, # noqa: F401
)
from app.schemas.plugin import (
PluginActionResponse, # noqa: F401
PluginInfo, # noqa: F401
PluginListResponse, # noqa: F401
PluginUninstallResponse, # noqa: F401
) )
from app.schemas.workflow import ( from app.schemas.workflow import (
WorkflowCreate, WorkflowUpdate, WorkflowResponse, WorkflowListResponse, AdvanceRequest, # noqa: F401
InstanceCreate, InstanceResponse, InstanceDetailResponse, InstanceListResponse, InstanceCreate, # noqa: F401
AdvanceRequest, StepHistoryResponse, WorkflowStep, InstanceDetailResponse, # noqa: F401
InstanceListResponse, # noqa: F401
InstanceResponse, # noqa: F401
StepHistoryResponse, # noqa: F401
WorkflowCreate, # noqa: F401
WorkflowListResponse, # noqa: F401
WorkflowResponse, # noqa: F401
WorkflowStep, # noqa: F401
WorkflowUpdate, # noqa: F401
) )
+7
View File
@@ -7,6 +7,7 @@ from pydantic import BaseModel, Field
class CopilotQueryRequest(BaseModel): class CopilotQueryRequest(BaseModel):
"""Natural language query to the AI copilot.""" """Natural language query to the AI copilot."""
query: str = Field(..., min_length=1, max_length=2000) query: str = Field(..., min_length=1, max_length=2000)
conversation_id: str | None = None conversation_id: str | None = None
context: dict = Field(default_factory=dict) context: dict = Field(default_factory=dict)
@@ -14,6 +15,7 @@ class CopilotQueryRequest(BaseModel):
class CopilotAction(BaseModel): class CopilotAction(BaseModel):
"""A proposed API action derived from NL input.""" """A proposed API action derived from NL input."""
method: str = Field(..., pattern="^(GET|POST|PATCH|DELETE)$") method: str = Field(..., pattern="^(GET|POST|PATCH|DELETE)$")
path: str = Field(..., min_length=1) path: str = Field(..., min_length=1)
body: dict | None = None body: dict | None = None
@@ -23,6 +25,7 @@ class CopilotAction(BaseModel):
class CopilotQueryResponse(BaseModel): class CopilotQueryResponse(BaseModel):
"""Response from copilot query — proposed actions for user confirmation.""" """Response from copilot query — proposed actions for user confirmation."""
conversation_id: str conversation_id: str
message: str message: str
proposed_actions: list[CopilotAction] = Field(default_factory=list) proposed_actions: list[CopilotAction] = Field(default_factory=list)
@@ -30,12 +33,14 @@ class CopilotQueryResponse(BaseModel):
class CopilotExecuteRequest(BaseModel): class CopilotExecuteRequest(BaseModel):
"""Execute a proposed action after user confirmation.""" """Execute a proposed action after user confirmation."""
conversation_id: str conversation_id: str
action: CopilotAction action: CopilotAction
class CopilotExecuteResponse(BaseModel): class CopilotExecuteResponse(BaseModel):
"""Result of executing a proposed action.""" """Result of executing a proposed action."""
conversation_id: str conversation_id: str
success: bool success: bool
status_code: int status_code: int
@@ -45,6 +50,7 @@ class CopilotExecuteResponse(BaseModel):
class CopilotMessageResponse(BaseModel): class CopilotMessageResponse(BaseModel):
"""A single message in conversation history.""" """A single message in conversation history."""
id: str id: str
role: str role: str
content: str content: str
@@ -56,6 +62,7 @@ class CopilotMessageResponse(BaseModel):
class CopilotHistoryResponse(BaseModel): class CopilotHistoryResponse(BaseModel):
"""Paginated conversation history.""" """Paginated conversation history."""
items: list[CopilotMessageResponse] items: list[CopilotMessageResponse]
total: int total: int
page: int page: int
+4 -1
View File
@@ -37,7 +37,9 @@ class PluginActionResponse(BaseModel):
status: str status: str
installed: bool = False installed: bool = False
active: bool = False active: bool = False
dropped_tables: list[str] = Field(default_factory=list, description="Tables dropped (uninstall with remove_data=true)") dropped_tables: list[str] = Field(
default_factory=list, description="Tables dropped (uninstall with remove_data=true)"
)
message: str = "" message: str = ""
@@ -49,6 +51,7 @@ class PluginUninstallResponse(PluginActionResponse):
class ManifestFieldDoc(BaseModel): class ManifestFieldDoc(BaseModel):
"""Documentation for a single manifest field.""" """Documentation for a single manifest field."""
type: str type: str
required: str required: str
description: str description: str
-2
View File
@@ -2,8 +2,6 @@
from __future__ import annotations from __future__ import annotations
from typing import Any
from pydantic import BaseModel, EmailStr, Field from pydantic import BaseModel, EmailStr, Field
+11
View File
@@ -7,6 +7,7 @@ from pydantic import BaseModel, Field
class WorkflowStep(BaseModel): class WorkflowStep(BaseModel):
"""A single step in a workflow definition.""" """A single step in a workflow definition."""
name: str = Field(..., min_length=1, max_length=200) name: str = Field(..., min_length=1, max_length=200)
type: str = Field(..., pattern="^(action|approval|notification|condition)$") type: str = Field(..., pattern="^(action|approval|notification|condition)$")
config: dict = Field(default_factory=dict) config: dict = Field(default_factory=dict)
@@ -15,6 +16,7 @@ class WorkflowStep(BaseModel):
class WorkflowCreate(BaseModel): class WorkflowCreate(BaseModel):
"""Create a new workflow definition.""" """Create a new workflow definition."""
name: str = Field(..., min_length=1, max_length=200) name: str = Field(..., min_length=1, max_length=200)
description: str | None = None description: str | None = None
trigger_event: str | None = None trigger_event: str | None = None
@@ -24,6 +26,7 @@ class WorkflowCreate(BaseModel):
class WorkflowUpdate(BaseModel): class WorkflowUpdate(BaseModel):
"""Update an existing workflow definition.""" """Update an existing workflow definition."""
name: str | None = Field(None, max_length=200) name: str | None = Field(None, max_length=200)
description: str | None = None description: str | None = None
trigger_event: str | None = None trigger_event: str | None = None
@@ -33,6 +36,7 @@ class WorkflowUpdate(BaseModel):
class WorkflowResponse(BaseModel): class WorkflowResponse(BaseModel):
"""Workflow definition response.""" """Workflow definition response."""
id: str id: str
name: str name: str
description: str | None = None description: str | None = None
@@ -46,6 +50,7 @@ class WorkflowResponse(BaseModel):
class WorkflowListResponse(BaseModel): class WorkflowListResponse(BaseModel):
"""Paginated workflow list.""" """Paginated workflow list."""
items: list[WorkflowResponse] items: list[WorkflowResponse]
total: int total: int
page: int page: int
@@ -54,12 +59,14 @@ class WorkflowListResponse(BaseModel):
class InstanceCreate(BaseModel): class InstanceCreate(BaseModel):
"""Create a new workflow instance.""" """Create a new workflow instance."""
context: dict = Field(default_factory=dict) context: dict = Field(default_factory=dict)
timeout_hours: int | None = None timeout_hours: int | None = None
class InstanceResponse(BaseModel): class InstanceResponse(BaseModel):
"""Workflow instance response with current state and history.""" """Workflow instance response with current state and history."""
id: str id: str
workflow_id: str workflow_id: str
status: str status: str
@@ -75,12 +82,14 @@ class InstanceResponse(BaseModel):
class InstanceDetailResponse(InstanceResponse): class InstanceDetailResponse(InstanceResponse):
"""Instance with step history.""" """Instance with step history."""
history: list[dict] = Field(default_factory=list) history: list[dict] = Field(default_factory=list)
workflow_name: str | None = None workflow_name: str | None = None
class InstanceListResponse(BaseModel): class InstanceListResponse(BaseModel):
"""Paginated instance list.""" """Paginated instance list."""
items: list[InstanceResponse] items: list[InstanceResponse]
total: int total: int
page: int page: int
@@ -89,12 +98,14 @@ class InstanceListResponse(BaseModel):
class AdvanceRequest(BaseModel): class AdvanceRequest(BaseModel):
"""Advance or reject a workflow instance step.""" """Advance or reject a workflow instance step."""
decision: str = Field(..., pattern="^(approve|reject)$") decision: str = Field(..., pattern="^(approve|reject)$")
comment: str | None = None comment: str | None = None
class StepHistoryResponse(BaseModel): class StepHistoryResponse(BaseModel):
"""Step history entry.""" """Step history entry."""
id: str id: str
instance_id: str instance_id: str
step_index: int step_index: int
+22 -10
View File
@@ -1,16 +1,28 @@
"""Service layer package.""" """Service layer package."""
from app.services.plugin_service import PluginService, get_plugin_service
from app.services.ai_copilot_service import ( from app.services.ai_copilot_service import (
process_query as copilot_process_query, execute_action as copilot_execute_action, # noqa: F401
execute_action as copilot_execute_action,
get_history as copilot_get_history,
) )
from app.services.ai_copilot_service import (
get_history as copilot_get_history, # noqa: F401
)
from app.services.ai_copilot_service import (
process_query as copilot_process_query, # noqa: F401
)
from app.services.plugin_service import PluginService, get_plugin_service # noqa: F401
from app.services.workflow_service import ( from app.services.workflow_service import (
create_workflow, list_workflows, get_workflow, advance_instance, # noqa: F401
update_workflow, delete_workflow, auto_reject_timeout, # noqa: F401
create_instance, list_instances, get_instance, cancel_instance, # noqa: F401
advance_instance, cancel_instance, check_timeout, # noqa: F401
check_timeout, auto_reject_timeout, create_instance, # noqa: F401
find_workflows_for_event, start_instance_for_event, create_workflow, # noqa: F401
delete_workflow, # noqa: F401
find_workflows_for_event, # noqa: F401
get_instance, # noqa: F401
get_workflow, # noqa: F401
list_instances, # noqa: F401
list_workflows, # noqa: F401
start_instance_for_event, # noqa: F401
update_workflow, # noqa: F401
) )
+32 -24
View File
@@ -3,15 +3,15 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime, timezone from datetime import UTC, datetime
from typing import Any from typing import Any
from sqlalchemy import select, func, desc, and_ from sqlalchemy import desc, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.ai.llm_client import get_llm_client from app.ai.llm_client import get_llm_client
from app.core.audit import log_audit from app.core.audit import log_audit
from app.core.auth import check_permission, filter_fields_by_permission from app.core.auth import check_permission
from app.models.ai_conversation import AIConversation, AIMessage from app.models.ai_conversation import AIConversation, AIMessage
from app.models.company import Company from app.models.company import Company
from app.models.contact import Contact from app.models.contact import Contact
@@ -136,7 +136,9 @@ async def process_query(
# Log to audit # Log to audit
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="query", action="query",
entity_type="ai_copilot", entity_type="ai_copilot",
entity_id=conversation.id, entity_id=conversation.id,
@@ -173,7 +175,7 @@ async def execute_action(
select(AIConversation).where( select(AIConversation).where(
AIConversation.id == conv_uuid, AIConversation.id == conv_uuid,
AIConversation.tenant_id == tenant_id, AIConversation.tenant_id == tenant_id,
) )
) )
conversation = result.scalar_one_or_none() conversation = result.scalar_one_or_none()
if conversation is None: if conversation is None:
@@ -222,7 +224,9 @@ async def execute_action(
# Log to audit # Log to audit
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="execute", action="execute",
entity_type="ai_copilot", entity_type="ai_copilot",
entity_id=conversation.id, entity_id=conversation.id,
@@ -266,17 +270,23 @@ async def get_history(
items: list[dict[str, Any]] = [] items: list[dict[str, Any]] = []
for conv in conversations: for conv in conversations:
# Get messages for each conversation # Get messages for each conversation
msg_q = select(AIMessage).where( msg_q = (
AIMessage.conversation_id == conv.id, select(AIMessage)
AIMessage.tenant_id == tenant_id, .where(
).order_by(AIMessage.message_index) AIMessage.conversation_id == conv.id,
AIMessage.tenant_id == tenant_id,
)
.order_by(AIMessage.message_index)
)
msg_result = await db.execute(msg_q) msg_result = await db.execute(msg_q)
messages = msg_result.scalars().all() messages = msg_result.scalars().all()
items.append({ items.append(
**_conversation_to_dict(conv), {
"messages": [_message_to_dict(m) for m in messages], **_conversation_to_dict(conv),
}) "messages": [_message_to_dict(m) for m in messages],
}
)
return { return {
"items": items, "items": items,
@@ -333,10 +343,7 @@ async def _exec_companies(
return { return {
"success": True, "success": True,
"status_code": 200, "status_code": 200,
"data": [ "data": [{"id": str(c.id), "name": c.name, "industry": c.industry} for c in companies],
{"id": str(c.id), "name": c.name, "industry": c.industry}
for c in companies
],
} }
elif method == "POST": elif method == "POST":
@@ -372,9 +379,13 @@ async def _exec_companies(
company = result.scalar_one_or_none() company = result.scalar_one_or_none()
if company is None: if company is None:
return {"error": "Company not found", "status_code": 404, "success": False} return {"error": "Company not found", "status_code": 404, "success": False}
company.deleted_at = datetime.now(timezone.utc) company.deleted_at = datetime.now(UTC)
await db.flush() await db.flush()
return {"success": True, "status_code": 200, "data": {"id": str(company.id), "deleted": True}} return {
"success": True,
"status_code": 200,
"data": {"id": str(company.id), "deleted": True},
}
elif method == "PATCH": elif method == "PATCH":
if not entity_id or entity_id == "{id}": if not entity_id or entity_id == "{id}":
@@ -424,10 +435,7 @@ async def _exec_contacts(
return { return {
"success": True, "success": True,
"status_code": 200, "status_code": 200,
"data": [ "data": [{"id": str(c.id), "name": c.name, "email": c.email} for c in contacts],
{"id": str(c.id), "name": c.name, "email": c.email}
for c in contacts
],
} }
elif method == "POST": elif method == "POST":
+29 -17
View File
@@ -2,9 +2,8 @@
from __future__ import annotations from __future__ import annotations
import hashlib
import uuid import uuid
from datetime import datetime, timedelta, timezone from datetime import UTC, datetime, timedelta
from typing import Any from typing import Any
import redis.asyncio as aioredis import redis.asyncio as aioredis
@@ -12,14 +11,19 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings from app.config import get_settings
from app.core.auth import (
create_session, get_session_data, hash_password, hash_token,
invalidate_session, update_session_tenant, verify_password,
)
from app.core.audit import log_audit from app.core.audit import log_audit
from app.core.auth import (
create_session,
get_session_data,
hash_password,
hash_token,
invalidate_session,
update_session_tenant,
verify_password,
)
from app.models.auth import PasswordResetToken from app.models.auth import PasswordResetToken
from app.models.user import User, UserTenant
from app.models.tenant import Tenant from app.models.tenant import Tenant
from app.models.user import User, UserTenant
class AuthService: class AuthService:
@@ -37,7 +41,7 @@ class AuthService:
Returns (session_id, csrf_token, user, tenant) or None. Returns (session_id, csrf_token, user, tenant) or None.
""" """
# Find user by email — need to check across tenants or use default tenant # Find user by email — need to check across tenants or use default tenant
q = select(User).where(User.email == email, User.is_active == True) q = select(User).where(User.email == email, User.is_active == True) # noqa: E712
result = await db.execute(q) result = await db.execute(q)
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
if user is None: if user is None:
@@ -53,7 +57,7 @@ class AuthService:
Tenant.slug == tenant_slug Tenant.slug == tenant_slug
) )
else: else:
ut_q = ut_q.where(UserTenant.is_default == True) ut_q = ut_q.where(UserTenant.is_default == True) # noqa: E712
ut_result = await db.execute(ut_q) ut_result = await db.execute(ut_q)
user_tenant = ut_result.scalar_one_or_none() user_tenant = ut_result.scalar_one_or_none()
@@ -75,7 +79,12 @@ class AuthService:
# Log the login in audit trail # Log the login in audit trail
await log_audit( await log_audit(
db, tenant.id, user.id, "login", "user", user.id, db,
tenant.id,
user.id,
"login",
"user",
user.id,
changes={"email": email}, changes={"email": email},
) )
@@ -153,7 +162,7 @@ class AuthService:
tenant_id: uuid.UUID | None = None, tenant_id: uuid.UUID | None = None,
) -> bool: ) -> bool:
"""Create a password reset token. Always returns True (no user enumeration).""" """Create a password reset token. Always returns True (no user enumeration)."""
q = select(User).where(User.email == email, User.is_active == True) q = select(User).where(User.email == email, User.is_active == True) # noqa: E712
result = await db.execute(q) result = await db.execute(q)
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
if user is None: if user is None:
@@ -166,14 +175,15 @@ class AuthService:
) )
prev_result = await db.execute(prev_q) prev_result = await db.execute(prev_q)
for prev_token in prev_result.scalars().all(): for prev_token in prev_result.scalars().all():
prev_token.used_at = datetime.now(timezone.utc) prev_token.used_at = datetime.now(UTC)
# Create new token # Create new token
import secrets import secrets
raw_token = secrets.token_urlsafe(32) raw_token = secrets.token_urlsafe(32)
token_hash = hash_token(raw_token) token_hash = hash_token(raw_token)
settings = get_settings() settings = get_settings()
expires_at = datetime.now(timezone.utc) + timedelta(hours=settings.password_reset_expiry_hours) expires_at = datetime.now(UTC) + timedelta(hours=settings.password_reset_expiry_hours)
reset_token = PasswordResetToken( reset_token = PasswordResetToken(
tenant_id=user.tenant_id, tenant_id=user.tenant_id,
@@ -206,7 +216,7 @@ class AuthService:
if reset_token is None: if reset_token is None:
return False return False
if reset_token.expires_at < datetime.now(timezone.utc): if reset_token.expires_at < datetime.now(UTC):
return False # Token expired return False # Token expired
# Get user # Get user
@@ -218,7 +228,7 @@ class AuthService:
# Update password # Update password
user.password_hash = hash_password(new_password) user.password_hash = hash_password(new_password)
reset_token.used_at = datetime.now(timezone.utc) reset_token.used_at = datetime.now(UTC)
await db.flush() await db.flush()
return True return True
@@ -229,6 +239,7 @@ class AuthService:
""" """
# This is a test helper — in production the token goes via email only # This is a test helper — in production the token goes via email only
import secrets import secrets
q = select(User).where(User.email == email) q = select(User).where(User.email == email)
result = await db.execute(q) result = await db.execute(q)
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
@@ -238,7 +249,7 @@ class AuthService:
raw_token = secrets.token_urlsafe(32) raw_token = secrets.token_urlsafe(32)
token_hash = hash_token(raw_token) token_hash = hash_token(raw_token)
settings = get_settings() settings = get_settings()
expires_at = datetime.now(timezone.utc) + timedelta(hours=settings.password_reset_expiry_hours) expires_at = datetime.now(UTC) + timedelta(hours=settings.password_reset_expiry_hours)
reset_token = PasswordResetToken( reset_token = PasswordResetToken(
tenant_id=user.tenant_id, tenant_id=user.tenant_id,
@@ -253,6 +264,7 @@ class AuthService:
async def create_expired_reset_token(self, db: AsyncSession, email: str) -> str | None: async def create_expired_reset_token(self, db: AsyncSession, email: str) -> str | None:
"""Create an already-expired reset token for testing.""" """Create an already-expired reset token for testing."""
import secrets import secrets
q = select(User).where(User.email == email) q = select(User).where(User.email == email)
result = await db.execute(q) result = await db.execute(q)
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
@@ -261,7 +273,7 @@ class AuthService:
raw_token = secrets.token_urlsafe(32) raw_token = secrets.token_urlsafe(32)
token_hash = hash_token(raw_token) token_hash = hash_token(raw_token)
expires_at = datetime.now(timezone.utc) - timedelta(hours=1) # Already expired expires_at = datetime.now(UTC) - timedelta(hours=1) # Already expired
reset_token = PasswordResetToken( reset_token = PasswordResetToken(
tenant_id=user.tenant_id, tenant_id=user.tenant_id,
+87 -39
View File
@@ -5,15 +5,15 @@ from __future__ import annotations
import csv import csv
import io import io
import uuid import uuid
from datetime import datetime, timezone from datetime import UTC, datetime
from typing import Any from typing import Any
from sqlalchemy import select, func, or_, desc, asc, delete from sqlalchemy import asc, delete, desc, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.audit import log_audit, log_deletion from app.core.audit import log_audit
from app.models.company import Company from app.models.company import Company
from app.models.contact import Contact, CompanyContact from app.models.contact import CompanyContact, Contact
def _company_to_dict(c: Company, include_contacts: bool = False) -> dict[str, Any]: def _company_to_dict(c: Company, include_contacts: bool = False) -> dict[str, Any]:
@@ -66,9 +66,7 @@ async def list_companies(
pattern = f"%{search}%" pattern = f"%{search}%"
base = base.where( base = base.where(
or_( or_(
Company.search_tsv.op("@@")( Company.search_tsv.op("@@")(func.plainto_tsquery("english", search)),
func.plainto_tsquery("english", search)
),
Company.name.ilike(pattern), Company.name.ilike(pattern),
Company.industry.ilike(pattern), Company.industry.ilike(pattern),
Company.description.ilike(pattern), Company.description.ilike(pattern),
@@ -131,16 +129,18 @@ async def get_company_detail(
contacts_result = await db.execute(contacts_q) contacts_result = await db.execute(contacts_q)
contacts_list = [] contacts_list = []
for contact, link in contacts_result.all(): for contact, link in contacts_result.all():
contacts_list.append({ contacts_list.append(
"id": str(contact.id), {
"first_name": contact.first_name, "id": str(contact.id),
"last_name": contact.last_name, "first_name": contact.first_name,
"email": contact.email, "last_name": contact.last_name,
"phone": contact.phone, "email": contact.email,
"position": contact.position, "phone": contact.phone,
"role_at_company": link.role_at_company, "position": contact.position,
"is_primary": link.is_primary, "role_at_company": link.role_at_company,
}) "is_primary": link.is_primary,
}
)
data["contacts"] = contacts_list data["contacts"] = contacts_list
return data return data
@@ -168,7 +168,12 @@ async def create_company(
await db.flush() await db.flush()
await db.refresh(company) await db.refresh(company)
await log_audit( await log_audit(
db, tenant_id, user_id, "create", "company", company.id, db,
tenant_id,
user_id,
"create",
"company",
company.id,
changes={"name": data["name"]}, changes={"name": data["name"]},
) )
return _company_to_dict(company) return _company_to_dict(company)
@@ -224,7 +229,7 @@ async def soft_delete_company(
if company is None: if company is None:
return False return False
company.deleted_at = datetime.now(timezone.utc) company.deleted_at = datetime.now(UTC)
company.updated_by = user_id company.updated_by = user_id
if cascade: if cascade:
@@ -238,7 +243,12 @@ async def soft_delete_company(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "delete", "company", company_id, db,
tenant_id,
user_id,
"delete",
"company",
company_id,
changes={"name": company.name, "cascade": cascade}, changes={"name": company.name, "cascade": cascade},
) )
return True return True
@@ -288,7 +298,12 @@ async def link_contact(
existing.is_primary = is_primary existing.is_primary = is_primary
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "link", "company_contact", existing.id, db,
tenant_id,
user_id,
"link",
"company_contact",
existing.id,
changes={"company_id": str(company_id), "contact_id": str(contact_id)}, changes={"company_id": str(company_id), "contact_id": str(contact_id)},
) )
return { return {
@@ -309,7 +324,12 @@ async def link_contact(
db.add(link) db.add(link)
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "link", "company_contact", link.id, db,
tenant_id,
user_id,
"link",
"company_contact",
link.id,
changes={"company_id": str(company_id), "contact_id": str(contact_id)}, changes={"company_id": str(company_id), "contact_id": str(contact_id)},
) )
return { return {
@@ -342,7 +362,12 @@ async def unlink_contact(
await db.delete(link) await db.delete(link)
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "unlink", "company_contact", link.id, db,
tenant_id,
user_id,
"unlink",
"company_contact",
link.id,
changes={"company_id": str(company_id), "contact_id": str(contact_id)}, changes={"company_id": str(company_id), "contact_id": str(contact_id)},
) )
return True return True
@@ -365,9 +390,7 @@ async def export_companies_csv(
pattern = f"%{search}%" pattern = f"%{search}%"
base = base.where( base = base.where(
or_( or_(
Company.search_tsv.op("@@")( Company.search_tsv.op("@@")(func.plainto_tsquery("english", search)),
func.plainto_tsquery("english", search)
),
Company.name.ilike(pattern), Company.name.ilike(pattern),
Company.industry.ilike(pattern), Company.industry.ilike(pattern),
Company.description.ilike(pattern), Company.description.ilike(pattern),
@@ -379,12 +402,22 @@ async def export_companies_csv(
output = io.StringIO() output = io.StringIO()
writer = csv.writer(output) writer = csv.writer(output)
writer.writerow(["id", "name", "account_number", "industry", "phone", "email", "website", "description"]) writer.writerow(
["id", "name", "account_number", "industry", "phone", "email", "website", "description"]
)
for c in companies: for c in companies:
writer.writerow([ writer.writerow(
str(c.id), c.name, c.account_number or "", c.industry or "", [
c.phone or "", c.email or "", c.website or "", c.description or "", str(c.id),
]) c.name,
c.account_number or "",
c.industry or "",
c.phone or "",
c.email or "",
c.website or "",
c.description or "",
]
)
return output.getvalue() return output.getvalue()
@@ -407,9 +440,7 @@ async def export_companies_xlsx(
pattern = f"%{search}%" pattern = f"%{search}%"
base = base.where( base = base.where(
or_( or_(
Company.search_tsv.op("@@")( Company.search_tsv.op("@@")(func.plainto_tsquery("english", search)),
func.plainto_tsquery("english", search)
),
Company.name.ilike(pattern), Company.name.ilike(pattern),
Company.industry.ilike(pattern), Company.industry.ilike(pattern),
Company.description.ilike(pattern), Company.description.ilike(pattern),
@@ -422,13 +453,30 @@ async def export_companies_xlsx(
wb = Workbook() wb = Workbook()
ws = wb.active ws = wb.active
ws.title = "Companies" ws.title = "Companies"
headers = ["id", "name", "account_number", "industry", "phone", "email", "website", "description"] headers = [
"id",
"name",
"account_number",
"industry",
"phone",
"email",
"website",
"description",
]
ws.append(headers) ws.append(headers)
for c in companies: for c in companies:
ws.append([ ws.append(
str(c.id), c.name, c.account_number or "", c.industry or "", [
c.phone or "", c.email or "", c.website or "", c.description or "", str(c.id),
]) c.name,
c.account_number or "",
c.industry or "",
c.phone or "",
c.email or "",
c.website or "",
c.description or "",
]
)
buf = io.BytesIO() buf = io.BytesIO()
wb.save(buf) wb.save(buf)
+47 -16
View File
@@ -3,15 +3,15 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime, timezone from datetime import UTC, datetime
from typing import Any from typing import Any
from sqlalchemy import select, func, desc, asc, delete from sqlalchemy import asc, delete, desc, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.audit import log_audit, log_deletion from app.core.audit import log_audit, log_deletion
from app.models.contact import Contact, CompanyContact
from app.models.company import Company from app.models.company import Company
from app.models.contact import CompanyContact, Contact
def _contact_to_dict(c: Contact, include_companies: bool = False) -> dict[str, Any]: def _contact_to_dict(c: Contact, include_companies: bool = False) -> dict[str, Any]:
@@ -113,13 +113,15 @@ async def get_contact_detail(
companies_result = await db.execute(companies_q) companies_result = await db.execute(companies_q)
companies_list = [] companies_list = []
for company, link in companies_result.all(): for company, link in companies_result.all():
companies_list.append({ companies_list.append(
"id": str(company.id), {
"name": company.name, "id": str(company.id),
"industry": company.industry, "name": company.name,
"role_at_company": link.role_at_company, "industry": company.industry,
"is_primary": link.is_primary, "role_at_company": link.role_at_company,
}) "is_primary": link.is_primary,
}
)
data["companies"] = companies_list data["companies"] = companies_list
return data return data
@@ -177,8 +179,17 @@ async def create_contact(
await db.refresh(contact) await db.refresh(contact)
await log_audit( await log_audit(
db, tenant_id, user_id, "create", "contact", contact.id, db,
changes={"first_name": data["first_name"], "last_name": data["last_name"], "linked_companies": linked_companies}, tenant_id,
user_id,
"create",
"contact",
contact.id,
changes={
"first_name": data["first_name"],
"last_name": data["last_name"],
"linked_companies": linked_companies,
},
) )
return _contact_to_dict(contact) return _contact_to_dict(contact)
@@ -202,7 +213,17 @@ async def update_contact(
return None return None
changes: dict[str, Any] = {} changes: dict[str, Any] = {}
for field in ("first_name", "last_name", "email", "phone", "mobile", "position", "department", "linkedin_url", "notes"): for field in (
"first_name",
"last_name",
"email",
"phone",
"mobile",
"position",
"department",
"linkedin_url",
"notes",
):
if field in data and data[field] is not None: if field in data and data[field] is not None:
old_val = getattr(contact, field) old_val = getattr(contact, field)
changes[field] = {"old": old_val, "new": data[field]} changes[field] = {"old": old_val, "new": data[field]}
@@ -232,12 +253,17 @@ async def soft_delete_contact(
if contact is None: if contact is None:
return False return False
contact.deleted_at = datetime.now(timezone.utc) contact.deleted_at = datetime.now(UTC)
contact.updated_by = user_id contact.updated_by = user_id
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "delete", "contact", contact_id, db,
tenant_id,
user_id,
"delete",
"contact",
contact_id,
changes={"first_name": contact.first_name, "last_name": contact.last_name}, changes={"first_name": contact.first_name, "last_name": contact.last_name},
) )
return True return True
@@ -276,6 +302,11 @@ async def gdpr_hard_delete_contact(
# Deletion log (immutable snapshot) # Deletion log (immutable snapshot)
await log_deletion( await log_deletion(
db, tenant_id, user_id, "contact", contact_id, snapshot, db,
tenant_id,
user_id,
"contact",
contact_id,
snapshot,
) )
return True return True
+35 -12
View File
@@ -16,7 +16,6 @@ from app.models.contact import Contact
from app.services.company_service import _company_to_dict from app.services.company_service import _company_to_dict
from app.services.contact_service import _contact_to_dict from app.services.contact_service import _contact_to_dict
# Expected CSV columns for each entity type # Expected CSV columns for each entity type
COMPANY_COLUMNS = ["name", "industry", "phone", "email", "website", "description"] COMPANY_COLUMNS = ["name", "industry", "phone", "email", "website", "description"]
CONTACT_COLUMNS = ["first_name", "last_name", "email", "phone", "mobile", "position", "department"] CONTACT_COLUMNS = ["first_name", "last_name", "email", "phone", "mobile", "position", "department"]
@@ -88,7 +87,12 @@ async def import_companies(
db.add(company) db.add(company)
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "import", "company", company.id, db,
tenant_id,
user_id,
"import",
"company",
company.id,
changes={"name": company.name}, changes={"name": company.name},
) )
created.append(_company_to_dict(company)) created.append(_company_to_dict(company))
@@ -154,7 +158,12 @@ async def import_contacts(
db.add(contact) db.add(contact)
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, "import", "contact", contact.id, db,
tenant_id,
user_id,
"import",
"contact",
contact.id,
changes={"first_name": contact.first_name, "last_name": contact.last_name}, changes={"first_name": contact.first_name, "last_name": contact.last_name},
) )
created.append(_contact_to_dict(contact)) created.append(_contact_to_dict(contact))
@@ -198,19 +207,33 @@ async def export_contacts_csv(
tenant_id: uuid.UUID, tenant_id: uuid.UUID,
) -> str: ) -> str:
"""Export contacts as CSV string.""" """Export contacts as CSV string."""
q = select(Contact).where( q = (
Contact.tenant_id == tenant_id, select(Contact)
Contact.deleted_at.is_(None), .where(
).order_by(Contact.last_name, Contact.first_name) Contact.tenant_id == tenant_id,
Contact.deleted_at.is_(None),
)
.order_by(Contact.last_name, Contact.first_name)
)
result = await db.execute(q) result = await db.execute(q)
contacts = result.scalars().all() contacts = result.scalars().all()
output = io.StringIO() output = io.StringIO()
writer = csv.writer(output) writer = csv.writer(output)
writer.writerow(["id", "first_name", "last_name", "email", "phone", "mobile", "position", "department"]) writer.writerow(
["id", "first_name", "last_name", "email", "phone", "mobile", "position", "department"]
)
for c in contacts: for c in contacts:
writer.writerow([ writer.writerow(
str(c.id), c.first_name, c.last_name, c.email or "", [
c.phone or "", c.mobile or "", c.position or "", c.department or "", str(c.id),
]) c.first_name,
c.last_name,
c.email or "",
c.phone or "",
c.mobile or "",
c.position or "",
c.department or "",
]
)
return output.getvalue() return output.getvalue()
+30 -13
View File
@@ -9,8 +9,8 @@ from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.audit import log_audit from app.core.audit import log_audit
from app.plugins.registry import PluginRegistry, get_registry
from app.plugins.migration_runner import MigrationValidationError from app.plugins.migration_runner import MigrationValidationError
from app.plugins.registry import PluginRegistry, get_registry
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -48,7 +48,9 @@ class PluginService:
record = await self._registry.install(db, name) record = await self._registry.install(db, name)
if tenant_id and user_id: if tenant_id and user_id:
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="plugin.install", action="plugin.install",
entity_type="plugin", entity_type="plugin",
entity_id=getattr(record, "id", None), entity_id=getattr(record, "id", None),
@@ -64,9 +66,9 @@ class PluginService:
"message": "Plugin installed successfully", "message": "Plugin installed successfully",
} }
except ValueError as exc: except ValueError as exc:
raise ValueError(str(exc)) raise ValueError(str(exc)) from None
except MigrationValidationError as exc: except MigrationValidationError as exc:
raise MigrationValidationError(str(exc)) raise MigrationValidationError(str(exc)) from None
async def activate_plugin( async def activate_plugin(
self, self,
@@ -84,7 +86,9 @@ class PluginService:
was_already_active = record.active and record.status == "active" was_already_active = record.active and record.status == "active"
if tenant_id and user_id: if tenant_id and user_id:
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="plugin.activate", action="plugin.activate",
entity_type="plugin", entity_type="plugin",
entity_id=getattr(record, "id", None), entity_id=getattr(record, "id", None),
@@ -97,10 +101,12 @@ class PluginService:
"status": record.status, "status": record.status,
"installed": record.installed, "installed": record.installed,
"active": record.active, "active": record.active,
"message": "Plugin is already active" if was_already_active else "Plugin activated successfully", "message": "Plugin is already active"
if was_already_active
else "Plugin activated successfully",
} }
except ValueError as exc: except ValueError as exc:
raise ValueError(str(exc)) raise ValueError(str(exc)) from None
async def deactivate_plugin( async def deactivate_plugin(
self, self,
@@ -118,7 +124,9 @@ class PluginService:
was_already_inactive = not record.active and record.status == "inactive" was_already_inactive = not record.active and record.status == "inactive"
if tenant_id and user_id: if tenant_id and user_id:
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="plugin.deactivate", action="plugin.deactivate",
entity_type="plugin", entity_type="plugin",
entity_id=getattr(record, "id", None), entity_id=getattr(record, "id", None),
@@ -131,10 +139,12 @@ class PluginService:
"status": record.status, "status": record.status,
"installed": record.installed, "installed": record.installed,
"active": record.active, "active": record.active,
"message": "Plugin is already inactive" if was_already_inactive else "Plugin deactivated successfully", "message": "Plugin is already inactive"
if was_already_inactive
else "Plugin deactivated successfully",
} }
except ValueError as exc: except ValueError as exc:
raise ValueError(str(exc)) raise ValueError(str(exc)) from None
async def uninstall_plugin( async def uninstall_plugin(
self, self,
@@ -153,10 +163,16 @@ class PluginService:
dropped_tables = getattr(record, "dropped_tables", []) dropped_tables = getattr(record, "dropped_tables", [])
if tenant_id and user_id: if tenant_id and user_id:
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="plugin.uninstall", action="plugin.uninstall",
entity_type="plugin", entity_type="plugin",
changes={"name": name, "remove_data": remove_data, "dropped_tables": dropped_tables}, changes={
"name": name,
"remove_data": remove_data,
"dropped_tables": dropped_tables,
},
) )
return { return {
"name": name, "name": name,
@@ -169,11 +185,12 @@ class PluginService:
"message": f"Plugin uninstalled{' and data tables dropped' if remove_data else ''}", "message": f"Plugin uninstalled{' and data tables dropped' if remove_data else ''}",
} }
except ValueError as exc: except ValueError as exc:
raise ValueError(str(exc)) raise ValueError(str(exc)) from None
def get_manifest_schema(self) -> dict[str, Any]: def get_manifest_schema(self) -> dict[str, Any]:
"""Return the manifest schema documentation.""" """Return the manifest schema documentation."""
from app.plugins.manifest import MANIFEST_SCHEMA_DOC from app.plugins.manifest import MANIFEST_SCHEMA_DOC
return MANIFEST_SCHEMA_DOC.model_dump() return MANIFEST_SCHEMA_DOC.model_dump()
+1 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import uuid import uuid
from typing import Any from typing import Any
from sqlalchemy import select, func, or_ from sqlalchemy import func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.auth import hash_password from app.core.auth import hash_password
+92 -46
View File
@@ -3,15 +3,15 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime, timezone, timedelta from datetime import UTC, datetime, timedelta
from typing import Any from typing import Any
from sqlalchemy import select, func, desc, and_, or_ from sqlalchemy import desc, func, select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.audit import log_audit from app.core.audit import log_audit
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
from app.models.notification import Notification from app.models.notification import Notification
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
def _safe_iso(dt) -> str | None: def _safe_iso(dt) -> str | None:
@@ -47,7 +47,12 @@ def _workflow_to_dict(w: Workflow) -> dict[str, Any]:
} }
def _instance_to_dict(i: WorkflowInstance, include_history: bool = False, history: list | None = None, workflow_name: str | None = None) -> dict[str, Any]: def _instance_to_dict(
i: WorkflowInstance,
include_history: bool = False,
history: list | None = None,
workflow_name: str | None = None,
) -> dict[str, Any]:
data = { data = {
"id": str(i.id), "id": str(i.id),
"workflow_id": str(i.workflow_id), "workflow_id": str(i.workflow_id),
@@ -107,6 +112,7 @@ async def _log_step_history(
# ─── Workflow CRUD ─── # ─── Workflow CRUD ───
async def create_workflow( async def create_workflow(
db: AsyncSession, db: AsyncSession,
tenant_id: uuid.UUID, tenant_id: uuid.UUID,
@@ -134,7 +140,9 @@ async def create_workflow(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="create", action="create",
entity_type="workflow", entity_type="workflow",
entity_id=workflow.id, entity_id=workflow.id,
@@ -229,7 +237,9 @@ async def update_workflow(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="update", action="update",
entity_type="workflow", entity_type="workflow",
entity_id=workflow.id, entity_id=workflow.id,
@@ -261,7 +271,9 @@ async def delete_workflow(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="delete", action="delete",
entity_type="workflow", entity_type="workflow",
entity_id=wf_uuid, entity_id=wf_uuid,
@@ -272,6 +284,7 @@ async def delete_workflow(
# ─── Instance Lifecycle ─── # ─── Instance Lifecycle ───
async def create_instance( async def create_instance(
db: AsyncSession, db: AsyncSession,
tenant_id: uuid.UUID, tenant_id: uuid.UUID,
@@ -294,7 +307,7 @@ async def create_instance(
timeout_at = None timeout_at = None
if timeout_hours: if timeout_hours:
timeout_at = datetime.now(timezone.utc) + timedelta(hours=timeout_hours) timeout_at = datetime.now(UTC) + timedelta(hours=timeout_hours)
instance = WorkflowInstance( instance = WorkflowInstance(
tenant_id=tenant_id, tenant_id=tenant_id,
@@ -314,7 +327,9 @@ async def create_instance(
steps = workflow.steps or [] steps = workflow.steps or []
first_step = steps[0] if steps else {} first_step = steps[0] if steps else {}
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=0, step_index=0,
step_type=first_step.get("type", "action"), step_type=first_step.get("type", "action"),
action="entered", action="entered",
@@ -323,7 +338,9 @@ async def create_instance(
) )
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="create", action="create",
entity_type="workflow_instance", entity_type="workflow_instance",
entity_id=instance.id, entity_id=instance.id,
@@ -383,16 +400,18 @@ async def get_instance(
return None return None
# Get workflow name # Get workflow name
wf_result = await db.execute( wf_result = await db.execute(select(Workflow.name).where(Workflow.id == instance.workflow_id))
select(Workflow.name).where(Workflow.id == instance.workflow_id)
)
workflow_name = wf_result.scalar_one_or_none() workflow_name = wf_result.scalar_one_or_none()
# Get step history # Get step history
hist_q = select(WorkflowStepHistory).where( hist_q = (
WorkflowStepHistory.instance_id == inst_uuid, select(WorkflowStepHistory)
WorkflowStepHistory.tenant_id == tenant_id, .where(
).order_by(WorkflowStepHistory.created_at) WorkflowStepHistory.instance_id == inst_uuid,
WorkflowStepHistory.tenant_id == tenant_id,
)
.order_by(WorkflowStepHistory.created_at)
)
hist_result = await db.execute(hist_q) hist_result = await db.execute(hist_q)
history = hist_result.scalars().all() history = hist_result.scalars().all()
@@ -430,12 +449,13 @@ async def advance_instance(
return None return None
if instance.status not in ("pending", "in_progress"): if instance.status not in ("pending", "in_progress"):
return {"error": f"Cannot advance instance with status {instance.status}", "status_code": 400} return {
"error": f"Cannot advance instance with status {instance.status}",
"status_code": 400,
}
# Get workflow definition # Get workflow definition
wf_result = await db.execute( wf_result = await db.execute(select(Workflow).where(Workflow.id == instance.workflow_id))
select(Workflow).where(Workflow.id == instance.workflow_id)
)
workflow = wf_result.scalar_one_or_none() workflow = wf_result.scalar_one_or_none()
if workflow is None: if workflow is None:
return {"error": "Workflow definition not found", "status_code": 404} return {"error": "Workflow definition not found", "status_code": 404}
@@ -447,9 +467,11 @@ async def advance_instance(
if decision == "reject": if decision == "reject":
instance.status = "rejected" instance.status = "rejected"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=current_idx, step_index=current_idx,
step_type=step_type, step_type=step_type,
action="rejected", action="rejected",
@@ -470,7 +492,9 @@ async def advance_instance(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="reject", action="reject",
entity_type="workflow_instance", entity_type="workflow_instance",
entity_id=instance.id, entity_id=instance.id,
@@ -486,7 +510,9 @@ async def advance_instance(
# Log approval # Log approval
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=current_idx, step_index=current_idx,
step_type=step_type, step_type=step_type,
action="approved", action="approved",
@@ -498,9 +524,11 @@ async def advance_instance(
if next_idx >= len(steps): if next_idx >= len(steps):
# Workflow complete # Workflow complete
instance.status = "completed" instance.status = "completed"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=current_idx, step_index=current_idx,
step_type="complete", step_type="complete",
action="completed", action="completed",
@@ -510,7 +538,9 @@ async def advance_instance(
instance.current_step_index = next_idx instance.current_step_index = next_idx
next_step = steps[next_idx] next_step = steps[next_idx]
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=next_idx, step_index=next_idx,
step_type=next_step.get("type", "action"), step_type=next_step.get("type", "action"),
action="entered", action="entered",
@@ -520,11 +550,17 @@ async def advance_instance(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="advance", action="advance",
entity_type="workflow_instance", entity_type="workflow_instance",
entity_id=instance.id, entity_id=instance.id,
changes={"decision": decision, "step": current_idx, "next_step": next_idx if next_idx < len(steps) else None}, changes={
"decision": decision,
"step": current_idx,
"next_step": next_idx if next_idx < len(steps) else None,
},
) )
return _instance_to_dict(instance) return _instance_to_dict(instance)
@@ -549,21 +585,26 @@ async def cancel_instance(
return None return None
if instance.status in ("completed", "rejected", "cancelled"): if instance.status in ("completed", "rejected", "cancelled"):
return {"error": f"Cannot cancel instance with status {instance.status}", "status_code": 400} return {
"error": f"Cannot cancel instance with status {instance.status}",
"status_code": 400,
}
instance.status = "cancelled" instance.status = "cancelled"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
# Get current step for history # Get current step for history
wf_result = await db.execute( wf_result = await db.execute(select(Workflow).where(Workflow.id == instance.workflow_id))
select(Workflow).where(Workflow.id == instance.workflow_id)
)
workflow = wf_result.scalar_one_or_none() workflow = wf_result.scalar_one_or_none()
steps = workflow.steps if workflow else [] steps = workflow.steps if workflow else []
current_step = steps[instance.current_step_index] if instance.current_step_index < len(steps) else {} current_step = (
steps[instance.current_step_index] if instance.current_step_index < len(steps) else {}
)
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type=current_step.get("type", "action"), step_type=current_step.get("type", "action"),
action="cancelled", action="cancelled",
@@ -573,7 +614,9 @@ async def cancel_instance(
await db.flush() await db.flush()
await log_audit( await log_audit(
db, tenant_id, user_id, db,
tenant_id,
user_id,
action="cancel", action="cancel",
entity_type="workflow_instance", entity_type="workflow_instance",
entity_id=instance.id, entity_id=instance.id,
@@ -591,7 +634,7 @@ async def check_timeout(instance: WorkflowInstance) -> bool:
return False return False
if instance.status not in ("pending", "in_progress"): if instance.status not in ("pending", "in_progress"):
return False return False
return datetime.now(timezone.utc) > instance.timeout_at return datetime.now(UTC) > instance.timeout_at
async def auto_reject_timeout( async def auto_reject_timeout(
@@ -601,17 +644,19 @@ async def auto_reject_timeout(
) -> dict[str, Any]: ) -> dict[str, Any]:
"""Auto-reject a timed-out instance.""" """Auto-reject a timed-out instance."""
instance.status = "rejected" instance.status = "rejected"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
wf_result = await db.execute( wf_result = await db.execute(select(Workflow).where(Workflow.id == instance.workflow_id))
select(Workflow).where(Workflow.id == instance.workflow_id)
)
workflow = wf_result.scalar_one_or_none() workflow = wf_result.scalar_one_or_none()
steps = workflow.steps if workflow else [] steps = workflow.steps if workflow else []
current_step = steps[instance.current_step_index] if instance.current_step_index < len(steps) else {} current_step = (
steps[instance.current_step_index] if instance.current_step_index < len(steps) else {}
)
await _log_step_history( await _log_step_history(
db, tenant_id, instance.id, db,
tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type=current_step.get("type", "approval"), step_type=current_step.get("type", "approval"),
action="auto_rejected", action="auto_rejected",
@@ -665,7 +710,8 @@ async def start_instance_for_event(
instances: list[dict[str, Any]] = [] instances: list[dict[str, Any]] = []
for wf in workflows: for wf in workflows:
inst = await create_instance( inst = await create_instance(
db, tenant_id, db,
tenant_id,
user_id or uuid.uuid4(), user_id or uuid.uuid4(),
str(wf.id), str(wf.id),
context=context, context=context,
+9 -3
View File
@@ -67,9 +67,10 @@ async def ensure_onboarding_workflow_exists(
If it doesn't exist yet, create it. Returns the workflow ID. If it doesn't exist yet, create it. Returns the workflow ID.
""" """
from app.models.workflow import Workflow
from sqlalchemy import select from sqlalchemy import select
from app.models.workflow import Workflow
result = await db.execute( result = await db.execute(
select(Workflow).where( select(Workflow).where(
Workflow.tenant_id == tenant_id, Workflow.tenant_id == tenant_id,
@@ -82,8 +83,11 @@ async def ensure_onboarding_workflow_exists(
return existing.id return existing.id
from app.services.workflow_service import create_workflow from app.services.workflow_service import create_workflow
wf_dict = await create_workflow( wf_dict = await create_workflow(
db, tenant_id, user_id, db,
tenant_id,
user_id,
get_onboarding_workflow_definition(), get_onboarding_workflow_definition(),
) )
return uuid.UUID(wf_dict["id"]) if wf_dict else None return uuid.UUID(wf_dict["id"]) if wf_dict else None
@@ -105,7 +109,9 @@ async def trigger_onboarding(
return None return None
instance = await create_instance( instance = await create_instance(
db, tenant_id, admin_user_id, db,
tenant_id,
admin_user_id,
str(wf_id), str(wf_id),
context={"new_user_id": str(new_user_id), "user_id": str(new_user_id)}, context={"new_user_id": str(new_user_id), "user_id": str(new_user_id)},
) )
+38 -20
View File
@@ -7,20 +7,20 @@ Integrates with the event bus for event-triggered workflows.
from __future__ import annotations from __future__ import annotations
import uuid
import logging import logging
from datetime import datetime, timezone import uuid
from datetime import UTC, datetime
from typing import Any from typing import Any
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.event_bus import get_event_bus from app.core.event_bus import get_event_bus
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
from app.models.notification import Notification from app.models.notification import Notification
from app.models.workflow import Workflow, WorkflowInstance
from app.services.workflow_service import ( from app.services.workflow_service import (
_log_step_history,
_instance_to_dict, _instance_to_dict,
_log_step_history,
create_instance, create_instance,
find_workflows_for_event, find_workflows_for_event,
) )
@@ -59,7 +59,7 @@ class WorkflowEngine:
steps = workflow.steps or [] steps = workflow.steps or []
if instance.current_step_index >= len(steps): if instance.current_step_index >= len(steps):
instance.status = "completed" instance.status = "completed"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
await self.db.flush() await self.db.flush()
return _instance_to_dict(instance) return _instance_to_dict(instance)
@@ -68,7 +68,9 @@ class WorkflowEngine:
# Log step entry # Log step entry
await _log_step_history( await _log_step_history(
self.db, self.tenant_id, instance.id, self.db,
self.tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type=step_type, step_type=step_type,
action="processing", action="processing",
@@ -103,7 +105,9 @@ class WorkflowEngine:
# Execute action based on type # Execute action based on type
if action_type == "create_notification": if action_type == "create_notification":
user_id = config.get("user_id") or (str(instance.initiated_by) if instance.initiated_by else None) user_id = config.get("user_id") or (
str(instance.initiated_by) if instance.initiated_by else None
)
if user_id: if user_id:
notification = Notification( notification = Notification(
tenant_id=self.tenant_id, tenant_id=self.tenant_id,
@@ -120,7 +124,9 @@ class WorkflowEngine:
# Log action executed # Log action executed
await _log_step_history( await _log_step_history(
self.db, self.tenant_id, instance.id, self.db,
self.tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type="action", step_type="action",
action="executed", action="executed",
@@ -131,7 +137,7 @@ class WorkflowEngine:
next_idx = instance.current_step_index + 1 next_idx = instance.current_step_index + 1
if next_idx >= len(steps): if next_idx >= len(steps):
instance.status = "completed" instance.status = "completed"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
else: else:
instance.current_step_index = next_idx instance.current_step_index = next_idx
instance.status = "in_progress" instance.status = "in_progress"
@@ -139,12 +145,12 @@ class WorkflowEngine:
await self.db.flush() await self.db.flush()
return _instance_to_dict(instance) return _instance_to_dict(instance)
async def _process_notification( async def _process_notification(self, instance: WorkflowInstance, step: dict) -> dict[str, Any]:
self, instance: WorkflowInstance, step: dict
) -> dict[str, Any]:
"""Process a notification step — sends notification and advances.""" """Process a notification step — sends notification and advances."""
config = step.get("config", {}) config = step.get("config", {})
user_id = config.get("user_id") or (str(instance.initiated_by) if instance.initiated_by else None) user_id = config.get("user_id") or (
str(instance.initiated_by) if instance.initiated_by else None
)
if user_id: if user_id:
notification = Notification( notification = Notification(
@@ -158,7 +164,9 @@ class WorkflowEngine:
await self.db.flush() await self.db.flush()
await _log_step_history( await _log_step_history(
self.db, self.tenant_id, instance.id, self.db,
self.tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type="notification", step_type="notification",
action="sent", action="sent",
@@ -174,7 +182,7 @@ class WorkflowEngine:
next_idx = instance.current_step_index + 1 next_idx = instance.current_step_index + 1
if next_idx >= len(steps): if next_idx >= len(steps):
instance.status = "completed" instance.status = "completed"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
else: else:
instance.current_step_index = next_idx instance.current_step_index = next_idx
@@ -211,10 +219,16 @@ class WorkflowEngine:
elif operator == "lt": elif operator == "lt":
condition_met = actual is not None and expected is not None and actual < expected condition_met = actual is not None and expected is not None and actual < expected
elif operator == "contains": elif operator == "contains":
condition_met = actual is not None and expected in actual if isinstance(actual, (str, list)) else False condition_met = (
actual is not None and expected in actual
if isinstance(actual, str | list)
else False
)
await _log_step_history( await _log_step_history(
self.db, self.tenant_id, instance.id, self.db,
self.tenant_id,
instance.id,
step_index=instance.current_step_index, step_index=instance.current_step_index,
step_type="condition", step_type="condition",
action="evaluated", action="evaluated",
@@ -230,7 +244,7 @@ class WorkflowEngine:
next_idx = instance.current_step_index + 1 next_idx = instance.current_step_index + 1
if next_idx >= len(steps): if next_idx >= len(steps):
instance.status = "completed" instance.status = "completed"
instance.completed_at = datetime.now(timezone.utc) instance.completed_at = datetime.now(UTC)
else: else:
instance.current_step_index = next_idx instance.current_step_index = next_idx
@@ -253,8 +267,11 @@ async def handle_event(
instances: list[dict[str, Any]] = [] instances: list[dict[str, Any]] = []
for wf in workflows: for wf in workflows:
inst = await create_instance( inst = await create_instance(
db, tenant_id, db,
uuid.UUID(payload.get("user_id", str(uuid.uuid4()))) if payload.get("user_id") else None or uuid.uuid4(), tenant_id,
uuid.UUID(payload.get("user_id", str(uuid.uuid4())))
if payload.get("user_id")
else None or uuid.uuid4(),
str(wf.id), str(wf.id),
context=payload, context=payload,
) )
@@ -274,6 +291,7 @@ def register_workflow_event_handlers() -> None:
async def _workflow_event_handler(payload: dict[str, Any]) -> None: async def _workflow_event_handler(payload: dict[str, Any]) -> None:
"""Handle events that may trigger workflows.""" """Handle events that may trigger workflows."""
from app.core.db import create_db_session from app.core.db import create_db_session
tenant_id_str = payload.get("tenant_id") tenant_id_str = payload.get("tenant_id")
event_name = payload.get("event", "") event_name = payload.get("event", "")
if not tenant_id_str or not event_name: if not tenant_id_str or not event_name:
+27 -19
View File
@@ -7,8 +7,6 @@ Auth helpers talk to the HTTP API (integration tests).
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import json
import uuid
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from typing import Any from typing import Any
@@ -24,23 +22,22 @@ from sqlalchemy.ext.asyncio import (
create_async_engine, create_async_engine,
) )
from app.config import get_settings
from app.core.db import Base, reset_engine_for_testing, close_engine
from app.core.auth import hash_password from app.core.auth import hash_password
from app.core.db import Base, close_engine, reset_engine_for_testing
from app.main import create_app
from app.models.ai_conversation import AIConversation, AIMessage # noqa: F401
from app.models.company import Company
from app.models.contact import CompanyContact, Contact # noqa: F401
from app.models.plugin import Plugin, PluginMigration # noqa: F401
from app.models.role import Role
from app.models.tenant import Tenant from app.models.tenant import Tenant
from app.models.user import User, UserTenant from app.models.user import User, UserTenant
from app.models.role import Role from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory # noqa: F401
from app.models.company import Company from app.plugins.builtins.entity_links.models import EntityLink # noqa: F401
from app.models.contact import Contact, CompanyContact from app.plugins.builtins.permissions.models import Permission, ShareLink # noqa: F401
from app.models.plugin import Plugin, PluginMigration from app.plugins.builtins.tags.models import Tag, TagAssignment # noqa: F401
from app.models.ai_conversation import AIConversation, AIMessage
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
# Import plugin models so Base.metadata.create_all includes their tables
from app.plugins.builtins.tags.models import Tag, TagAssignment
from app.plugins.builtins.permissions.models import Permission, ShareLink
from app.plugins.builtins.entity_links.models import EntityLink
from app.main import create_app
# Import plugin models so Base.metadata.create_all includes their tables
TEST_DB_URL = "postgresql+asyncpg://leocrm:leocrm@localhost:5432/leocrm_test" TEST_DB_URL = "postgresql+asyncpg://leocrm:leocrm@localhost:5432/leocrm_test"
@@ -50,6 +47,7 @@ def _get_sync_engine():
Uses postgres superuser because leocrm user doesn't own the public schema. Uses postgres superuser because leocrm user doesn't own the public schema.
""" """
from sqlalchemy import create_engine from sqlalchemy import create_engine
return create_engine( return create_engine(
"postgresql+psycopg2://postgres@localhost:5432/leocrm_test", "postgresql+psycopg2://postgres@localhost:5432/leocrm_test",
echo=False, echo=False,
@@ -94,7 +92,11 @@ def clean_tables(db_setup):
sync_eng = _get_sync_engine() sync_eng = _get_sync_engine()
with sync_eng.connect() as conn: with sync_eng.connect() as conn:
# TRUNCATE all tables with CASCADE — fast and reliable isolation # TRUNCATE all tables with CASCADE — fast and reliable isolation
conn.execute(text("TRUNCATE TABLE entity_links, share_links, permissions, tag_assignments, tags, workflow_step_history, workflow_instances, workflows, ai_messages, ai_conversations, plugin_migrations, plugins, company_contacts, contacts, api_tokens, password_reset_tokens, notifications, deletion_log, audit_log, sessions, roles, companies, user_tenants, users, tenants CASCADE;")) conn.execute(
text(
"TRUNCATE TABLE entity_links, share_links, permissions, tag_assignments, tags, workflow_step_history, workflow_instances, workflows, ai_messages, ai_conversations, plugin_migrations, plugins, company_contacts, contacts, api_tokens, password_reset_tokens, notifications, deletion_log, audit_log, sessions, roles, companies, user_tenants, users, tenants CASCADE;"
)
)
conn.commit() conn.commit()
sync_eng.dispose() sync_eng.dispose()
yield yield
@@ -125,7 +127,9 @@ async def session_factory(engine: AsyncEngine) -> async_sessionmaker[AsyncSessio
@pytest_asyncio.fixture @pytest_asyncio.fixture
async def db_session(session_factory: async_sessionmaker[AsyncSession]) -> AsyncGenerator[AsyncSession, None]: async def db_session(
session_factory: async_sessionmaker[AsyncSession],
) -> AsyncGenerator[AsyncSession, None]:
"""Database session for direct DB operations in tests.""" """Database session for direct DB operations in tests."""
async with session_factory() as session: async with session_factory() as session:
yield session yield session
@@ -260,7 +264,9 @@ async def seed_tenant_and_users(db: AsyncSession) -> dict[str, Any]:
} }
async def login_client(client: AsyncClient, email: str, password: str = "TestPass123!") -> dict[str, str]: async def login_client(
client: AsyncClient, email: str, password: str = "TestPass123!"
) -> dict[str, str]:
"""Login via HTTP API and return cookies dict.""" """Login via HTTP API and return cookies dict."""
resp = await client.post( resp = await client.post(
"/api/v1/auth/login", "/api/v1/auth/login",
@@ -271,7 +277,9 @@ async def login_client(client: AsyncClient, email: str, password: str = "TestPas
return dict(resp.cookies) return dict(resp.cookies)
async def get_auth_client(client: AsyncClient, email: str, password: str = "TestPass123!") -> AsyncClient: async def get_auth_client(
client: AsyncClient, email: str, password: str = "TestPass123!"
) -> AsyncClient:
"""Return a client that's logged in.""" """Return a client that's logged in."""
await login_client(client, email, password) await login_client(client, email, password)
return client return client
+247 -47
View File
@@ -2,10 +2,12 @@
from __future__ import annotations from __future__ import annotations
from datetime import UTC
import pytest import pytest
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -68,7 +70,10 @@ async def test_ac3_copilot_execute_blocked_by_rbac(client: AsyncClient, db_sessi
# Query for a delete action # Query for a delete action
query_resp = await client.post( query_resp = await client.post(
"/api/v1/ai/copilot/query", "/api/v1/ai/copilot/query",
json={"query": "Delete company", "context": {"entity_id": "00000000-0000-0000-0000-000000000000"}}, json={
"query": "Delete company",
"context": {"entity_id": "00000000-0000-0000-0000-000000000000"},
},
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
assert query_resp.status_code == 200 assert query_resp.status_code == 200
@@ -119,6 +124,7 @@ async def test_ac4_copilot_history_paginated(client: AsyncClient, db_session):
async def test_ac5_copilot_action_logged_in_audit(client: AsyncClient, db_session): async def test_ac5_copilot_action_logged_in_audit(client: AsyncClient, db_session):
"""AC5: Copilot action logged in audit_log with entity_type=ai_copilot.""" """AC5: Copilot action logged in audit_log with entity_type=ai_copilot."""
from sqlalchemy import select from sqlalchemy import select
from app.models.audit import AuditLog from app.models.audit import AuditLog
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
@@ -133,9 +139,7 @@ async def test_ac5_copilot_action_logged_in_audit(client: AsyncClient, db_sessio
assert resp.status_code == 200 assert resp.status_code == 200
# Check audit log # Check audit log
result = await db_session.execute( result = await db_session.execute(select(AuditLog).where(AuditLog.entity_type == "ai_copilot"))
select(AuditLog).where(AuditLog.entity_type == "ai_copilot")
)
logs = result.scalars().all() logs = result.scalars().all()
assert len(logs) >= 1 assert len(logs) >= 1
assert logs[0].action == "query" assert logs[0].action == "query"
@@ -158,12 +162,14 @@ async def test_ac6_copilot_tenant_isolation(client: AsyncClient, db_session):
conv_id_a = query_resp.json()["conversation_id"] conv_id_a = query_resp.json()["conversation_id"]
# Login as tenant B admin (different cookie jar) # Login as tenant B admin (different cookie jar)
client2 = AsyncClient(transport=ASGITransport(app=client._transport.app), base_url="http://test") AsyncClient(transport=ASGITransport(app=client._transport.app), base_url="http://test")
# Need to use the same app — just re-login with a fresh client # Need to use the same app — just re-login with a fresh client
# Actually we need a new client without tenant A cookies # Actually we need a new client without tenant A cookies
from httpx import AsyncClient as AC from httpx import AsyncClient as AC # noqa: N817
# Use the same app instance # Use the same app instance
import app.main import app.main
app_instance = app.main.app app_instance = app.main.app
async with AC(transport=ASGITransport(app=app_instance), base_url="http://test") as client_b: async with AC(transport=ASGITransport(app=app_instance), base_url="http://test") as client_b:
@@ -174,7 +180,12 @@ async def test_ac6_copilot_tenant_isolation(client: AsyncClient, db_session):
"/api/v1/ai/copilot/execute", "/api/v1/ai/copilot/execute",
json={ json={
"conversation_id": conv_id_a, "conversation_id": conv_id_a,
"action": {"method": "GET", "path": "/api/v1/companies", "description": "List", "confidence": 0.9}, "action": {
"method": "GET",
"path": "/api/v1/companies",
"description": "List",
"confidence": 0.9,
},
}, },
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
@@ -254,9 +265,11 @@ async def test_copilot_unauthenticated(client: AsyncClient, db_session):
# ─── ActionMapper Unit Tests ─── # ─── ActionMapper Unit Tests ───
def test_action_mapper_create_company(): def test_action_mapper_create_company():
"""ActionMapper: 'create company named X' → POST /api/v1/companies.""" """ActionMapper: 'create company named X' → POST /api/v1/companies."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Create a company named Acme Corp") actions = map_query_to_actions("Create a company named Acme Corp")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "POST" assert actions[0]["method"] == "POST"
@@ -268,6 +281,7 @@ def test_action_mapper_create_company():
def test_action_mapper_create_company_no_name(): def test_action_mapper_create_company_no_name():
"""ActionMapper: 'create company' without name → default 'New Company'.""" """ActionMapper: 'create company' without name → default 'New Company'."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Add a new company") actions = map_query_to_actions("Add a new company")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["body"]["name"] == "New Company" assert actions[0]["body"]["name"] == "New Company"
@@ -276,6 +290,7 @@ def test_action_mapper_create_company_no_name():
def test_action_mapper_delete_company_with_context(): def test_action_mapper_delete_company_with_context():
"""ActionMapper: 'delete company' with entity_id in context → DELETE with specific ID.""" """ActionMapper: 'delete company' with entity_id in context → DELETE with specific ID."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
test_id = "12345678-1234-1234-1234-123456789abc" test_id = "12345678-1234-1234-1234-123456789abc"
actions = map_query_to_actions("Delete company", context={"entity_id": test_id}) actions = map_query_to_actions("Delete company", context={"entity_id": test_id})
assert len(actions) == 1 assert len(actions) == 1
@@ -287,6 +302,7 @@ def test_action_mapper_delete_company_with_context():
def test_action_mapper_delete_company_no_context(): def test_action_mapper_delete_company_no_context():
"""ActionMapper: 'delete company' without context → DELETE with {id} placeholder.""" """ActionMapper: 'delete company' without context → DELETE with {id} placeholder."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Remove company") actions = map_query_to_actions("Remove company")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "DELETE" assert actions[0]["method"] == "DELETE"
@@ -297,6 +313,7 @@ def test_action_mapper_delete_company_no_context():
def test_action_mapper_update_company(): def test_action_mapper_update_company():
"""ActionMapper: 'update company' → PATCH with extracted fields.""" """ActionMapper: 'update company' → PATCH with extracted fields."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Update company industry to Tech, name to FooBar") actions = map_query_to_actions("Update company industry to Tech, name to FooBar")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "PATCH" assert actions[0]["method"] == "PATCH"
@@ -307,6 +324,7 @@ def test_action_mapper_update_company():
def test_action_mapper_update_company_with_context(): def test_action_mapper_update_company_with_context():
"""ActionMapper: 'update company' with entity_id → PATCH with specific path.""" """ActionMapper: 'update company' with entity_id → PATCH with specific path."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
test_id = "12345678-1234-1234-1234-123456789abc" test_id = "12345678-1234-1234-1234-123456789abc"
actions = map_query_to_actions("Edit company", context={"company_id": test_id}) actions = map_query_to_actions("Edit company", context={"company_id": test_id})
assert len(actions) == 1 assert len(actions) == 1
@@ -316,6 +334,7 @@ def test_action_mapper_update_company_with_context():
def test_action_mapper_list_companies(): def test_action_mapper_list_companies():
"""ActionMapper: 'list companies' → GET /api/v1/companies.""" """ActionMapper: 'list companies' → GET /api/v1/companies."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Show all compan") actions = map_query_to_actions("Show all compan")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "GET" assert actions[0]["method"] == "GET"
@@ -325,6 +344,7 @@ def test_action_mapper_list_companies():
def test_action_mapper_list_companies_with_search(): def test_action_mapper_list_companies_with_search():
"""ActionMapper: 'find companies named X' → GET with search description.""" """ActionMapper: 'find companies named X' → GET with search description."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Find compan named Acme") actions = map_query_to_actions("Find compan named Acme")
assert len(actions) == 1 assert len(actions) == 1
assert "Acme" in actions[0]["description"] assert "Acme" in actions[0]["description"]
@@ -333,6 +353,7 @@ def test_action_mapper_list_companies_with_search():
def test_action_mapper_create_contact(): def test_action_mapper_create_contact():
"""ActionMapper: 'create contact named X' → POST /api/v1/contacts.""" """ActionMapper: 'create contact named X' → POST /api/v1/contacts."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Create a contact named John Doe") actions = map_query_to_actions("Create a contact named John Doe")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "POST" assert actions[0]["method"] == "POST"
@@ -343,6 +364,7 @@ def test_action_mapper_create_contact():
def test_action_mapper_list_contacts(): def test_action_mapper_list_contacts():
"""ActionMapper: 'list contacts' → GET /api/v1/contacts.""" """ActionMapper: 'list contacts' → GET /api/v1/contacts."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Show all contact") actions = map_query_to_actions("Show all contact")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "GET" assert actions[0]["method"] == "GET"
@@ -352,6 +374,7 @@ def test_action_mapper_list_contacts():
def test_action_mapper_list_workflows(): def test_action_mapper_list_workflows():
"""ActionMapper: 'list workflows' → GET /api/v1/workflows.""" """ActionMapper: 'list workflows' → GET /api/v1/workflows."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Show all workflow") actions = map_query_to_actions("Show all workflow")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "GET" assert actions[0]["method"] == "GET"
@@ -361,6 +384,7 @@ def test_action_mapper_list_workflows():
def test_action_mapper_create_workflow(): def test_action_mapper_create_workflow():
"""ActionMapper: 'create workflow named X' → POST /api/v1/workflows.""" """ActionMapper: 'create workflow named X' → POST /api/v1/workflows."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Create a new workflow named Approval Process") actions = map_query_to_actions("Create a new workflow named Approval Process")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "POST" assert actions[0]["method"] == "POST"
@@ -371,6 +395,7 @@ def test_action_mapper_create_workflow():
def test_action_mapper_help_intent(): def test_action_mapper_help_intent():
"""ActionMapper: 'help' query returns demo action with low confidence.""" """ActionMapper: 'help' query returns demo action with low confidence."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("What can you do? Help me please") actions = map_query_to_actions("What can you do? Help me please")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["confidence"] == 0.3 assert actions[0]["confidence"] == 0.3
@@ -379,6 +404,7 @@ def test_action_mapper_help_intent():
def test_action_mapper_unknown_query(): def test_action_mapper_unknown_query():
"""ActionMapper: unrecognized query returns empty actions list.""" """ActionMapper: unrecognized query returns empty actions list."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("xyzzy flonk") actions = map_query_to_actions("xyzzy flonk")
assert actions == [] assert actions == []
@@ -386,6 +412,7 @@ def test_action_mapper_unknown_query():
def test_action_mapper_update_company_phone_email(): def test_action_mapper_update_company_phone_email():
"""ActionMapper: update company with phone and email fields.""" """ActionMapper: update company with phone and email fields."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Update company phone to 123456, email to test@examplecom") actions = map_query_to_actions("Update company phone to 123456, email to test@examplecom")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["body"]["phone"] == "123456" assert actions[0]["body"]["phone"] == "123456"
@@ -395,6 +422,7 @@ def test_action_mapper_update_company_phone_email():
def test_action_mapper_update_company_no_fields(): def test_action_mapper_update_company_no_fields():
"""ActionMapper: update company without explicit fields → default name.""" """ActionMapper: update company without explicit fields → default name."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Modify company details") actions = map_query_to_actions("Modify company details")
assert len(actions) == 1 assert len(actions) == 1
assert "name" in actions[0]["body"] assert "name" in actions[0]["body"]
@@ -403,6 +431,7 @@ def test_action_mapper_update_company_no_fields():
def test_action_mapper_company_list_all_pattern(): def test_action_mapper_company_list_all_pattern():
"""ActionMapper: 'companies all' matches list_company2 pattern.""" """ActionMapper: 'companies all' matches list_company2 pattern."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("Show me companies all") actions = map_query_to_actions("Show me companies all")
assert len(actions) == 1 assert len(actions) == 1
assert actions[0]["method"] == "GET" assert actions[0]["method"] == "GET"
@@ -411,16 +440,19 @@ def test_action_mapper_company_list_all_pattern():
def test_action_mapper_empty_query(): def test_action_mapper_empty_query():
"""ActionMapper: empty query returns no actions.""" """ActionMapper: empty query returns no actions."""
from app.ai.action_mapper import map_query_to_actions from app.ai.action_mapper import map_query_to_actions
actions = map_query_to_actions("") actions = map_query_to_actions("")
assert actions == [] assert actions == []
# ─── LLMClient Unit Tests ─── # ─── LLMClient Unit Tests ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_llm_client_mock_mode_generate(): async def test_llm_client_mock_mode_generate():
"""LLMClient: mock mode generates actions from keyword matching.""" """LLMClient: mock mode generates actions from keyword matching."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
assert client.is_mock is True assert client.is_mock is True
@@ -434,6 +466,7 @@ async def test_llm_client_mock_mode_generate():
async def test_llm_client_mock_mode_no_actions(): async def test_llm_client_mock_mode_no_actions():
"""LLMClient: mock mode with unrecognized query returns empty actions.""" """LLMClient: mock mode with unrecognized query returns empty actions."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
response = await client.generate("xyzzy flonk") response = await client.generate("xyzzy flonk")
@@ -445,6 +478,7 @@ async def test_llm_client_mock_mode_no_actions():
async def test_llm_client_mock_mode_with_context(): async def test_llm_client_mock_mode_with_context():
"""LLMClient: mock mode passes context to action mapper.""" """LLMClient: mock mode passes context to action mapper."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
response = await client.generate("Delete company", context={"entity_id": "abc123"}) response = await client.generate("Delete company", context={"entity_id": "abc123"})
@@ -455,6 +489,7 @@ async def test_llm_client_mock_mode_with_context():
def test_llm_client_api_mode_init(): def test_llm_client_api_mode_init():
"""LLMClient: with model and api_key set, is_mock is False.""" """LLMClient: with model and api_key set, is_mock is False."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model="gpt-4", api_key="test-key") client = LLMClient(model="gpt-4", api_key="test-key")
assert client.is_mock is False assert client.is_mock is False
assert client.model == "gpt-4" assert client.model == "gpt-4"
@@ -464,6 +499,7 @@ def test_llm_client_api_mode_init():
def test_llm_client_api_base_default(): def test_llm_client_api_base_default():
"""LLMClient: default api_base is OpenAI URL.""" """LLMClient: default api_base is OpenAI URL."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
assert client.api_base == "https://api.openai.com/v1" assert client.api_base == "https://api.openai.com/v1"
@@ -471,6 +507,7 @@ def test_llm_client_api_base_default():
def test_llm_client_to_dict(): def test_llm_client_to_dict():
"""LLMResponse: to_dict returns structured response.""" """LLMResponse: to_dict returns structured response."""
from app.ai.llm_client import LLMResponse from app.ai.llm_client import LLMResponse
resp = LLMResponse(message="Hello", proposed_actions=[{"method": "GET"}], confidence=0.9) resp = LLMResponse(message="Hello", proposed_actions=[{"method": "GET"}], confidence=0.9)
d = resp.to_dict() d = resp.to_dict()
assert d["message"] == "Hello" assert d["message"] == "Hello"
@@ -480,14 +517,18 @@ def test_llm_client_to_dict():
def test_llm_client_parse_valid_json(): def test_llm_client_parse_valid_json():
"""LLMClient: _parse_llm_response with valid JSON returns structured response.""" """LLMClient: _parse_llm_response with valid JSON returns structured response."""
from app.ai.llm_client import LLMClient
import json import json
from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
content = json.dumps({ content = json.dumps(
"message": "Here are actions", {
"proposed_actions": [{"method": "GET", "path": "/api/v1/companies"}], "message": "Here are actions",
"confidence": 0.95, "proposed_actions": [{"method": "GET", "path": "/api/v1/companies"}],
}) "confidence": 0.95,
}
)
response = client._parse_llm_response(content) response = client._parse_llm_response(content)
assert response.message == "Here are actions" assert response.message == "Here are actions"
assert len(response.proposed_actions) == 1 assert len(response.proposed_actions) == 1
@@ -497,6 +538,7 @@ def test_llm_client_parse_valid_json():
def test_llm_client_parse_invalid_json(): def test_llm_client_parse_invalid_json():
"""LLMClient: _parse_llm_response with invalid JSON returns fallback response.""" """LLMClient: _parse_llm_response with invalid JSON returns fallback response."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
response = client._parse_llm_response("not valid json at all") response = client._parse_llm_response("not valid json at all")
assert response.proposed_actions == [] assert response.proposed_actions == []
@@ -506,6 +548,7 @@ def test_llm_client_parse_invalid_json():
def test_llm_client_build_system_prompt(): def test_llm_client_build_system_prompt():
"""LLMClient: _build_system_prompt contains API endpoints and context.""" """LLMClient: _build_system_prompt contains API endpoints and context."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
prompt = client._build_system_prompt({"page": "companies"}) prompt = client._build_system_prompt({"page": "companies"})
assert "LeoCRM" in prompt assert "LeoCRM" in prompt
@@ -516,6 +559,7 @@ def test_llm_client_build_system_prompt():
def test_llm_client_build_system_prompt_empty_context(): def test_llm_client_build_system_prompt_empty_context():
"""LLMClient: _build_system_prompt with no context uses empty dict.""" """LLMClient: _build_system_prompt with no context uses empty dict."""
from app.ai.llm_client import LLMClient from app.ai.llm_client import LLMClient
client = LLMClient(model=None, api_key=None) client = LLMClient(model=None, api_key=None)
prompt = client._build_system_prompt({}) prompt = client._build_system_prompt({})
assert "LeoCRM" in prompt assert "LeoCRM" in prompt
@@ -523,7 +567,8 @@ def test_llm_client_build_system_prompt_empty_context():
def test_llm_client_get_and_reset(): def test_llm_client_get_and_reset():
"""LLMClient: get_llm_client returns singleton, reset clears it.""" """LLMClient: get_llm_client returns singleton, reset clears it."""
from app.ai.llm_client import get_llm_client, reset_llm_client, LLMClient from app.ai.llm_client import get_llm_client, reset_llm_client
reset_llm_client() reset_llm_client()
client1 = get_llm_client() client1 = get_llm_client()
client2 = get_llm_client() client2 = get_llm_client()
@@ -535,10 +580,12 @@ def test_llm_client_get_and_reset():
# ─── AI Copilot Service Unit Tests ─── # ─── AI Copilot Service Unit Tests ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_service_process_query_new_conversation(db_session): async def test_service_process_query_new_conversation(db_session):
"""Service: process_query creates new conversation and returns proposed actions.""" """Service: process_query creates new conversation and returns proposed actions."""
from app.services.ai_copilot_service import process_query from app.services.ai_copilot_service import process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -553,6 +600,7 @@ async def test_service_process_query_new_conversation(db_session):
async def test_service_process_query_existing_conversation(db_session): async def test_service_process_query_existing_conversation(db_session):
"""Service: process_query with existing conversation_id appends to conversation.""" """Service: process_query with existing conversation_id appends to conversation."""
from app.services.ai_copilot_service import process_query from app.services.ai_copilot_service import process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -563,7 +611,9 @@ async def test_service_process_query_existing_conversation(db_session):
# Second query with same conversation_id # Second query with same conversation_id
result2 = await process_query( result2 = await process_query(
db_session, tenant_id, admin_id, db_session,
tenant_id,
admin_id,
"Create a company named FooBar", "Create a company named FooBar",
conversation_id=conv_id, conversation_id=conv_id,
) )
@@ -574,12 +624,16 @@ async def test_service_process_query_existing_conversation(db_session):
async def test_service_process_query_invalid_conversation(db_session): async def test_service_process_query_invalid_conversation(db_session):
"""Service: process_query with invalid conversation_id returns 404 error.""" """Service: process_query with invalid conversation_id returns 404 error."""
from app.services.ai_copilot_service import process_query from app.services.ai_copilot_service import process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
result = await process_query( result = await process_query(
db_session, tenant_id, admin_id, "List companies", db_session,
tenant_id,
admin_id,
"List companies",
conversation_id="00000000-0000-0000-0000-000000000000", conversation_id="00000000-0000-0000-0000-000000000000",
) )
assert result["error"] == "Conversation not found" assert result["error"] == "Conversation not found"
@@ -590,6 +644,7 @@ async def test_service_process_query_invalid_conversation(db_session):
async def test_service_process_query_empty_query(db_session): async def test_service_process_query_empty_query(db_session):
"""Service: process_query with empty query creates conversation with 'Untitled' title.""" """Service: process_query with empty query creates conversation with 'Untitled' title."""
from app.services.ai_copilot_service import process_query from app.services.ai_copilot_service import process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -602,17 +657,23 @@ async def test_service_process_query_empty_query(db_session):
async def test_service_execute_action_companies_get(db_session): async def test_service_execute_action_companies_get(db_session):
"""Service: execute_action with GET /api/v1/companies returns list.""" """Service: execute_action with GET /api/v1/companies returns list."""
from app.services.ai_copilot_service import execute_action from app.services.ai_copilot_service import execute_action
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
# First create a conversation # First create a conversation
from app.services.ai_copilot_service import process_query from app.services.ai_copilot_service import process_query
query_result = await process_query(db_session, tenant_id, admin_id, "List companies") query_result = await process_query(db_session, tenant_id, admin_id, "List companies")
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "GET", "path": "/api/v1/companies", "body": None}, {"method": "GET", "path": "/api/v1/companies", "body": None},
) )
assert result["success"] is True assert result["success"] is True
@@ -624,6 +685,7 @@ async def test_service_execute_action_companies_get(db_session):
async def test_service_execute_action_companies_post(db_session): async def test_service_execute_action_companies_post(db_session):
"""Service: execute_action with POST /api/v1/companies creates company.""" """Service: execute_action with POST /api/v1/companies creates company."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -632,8 +694,16 @@ async def test_service_execute_action_companies_post(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
{"method": "POST", "path": "/api/v1/companies", "body": {"name": "NewCo", "industry": "Tech"}}, tenant_id,
admin_id,
"admin",
conv_id,
{
"method": "POST",
"path": "/api/v1/companies",
"body": {"name": "NewCo", "industry": "Tech"},
},
) )
assert result["success"] is True assert result["success"] is True
assert result["status_code"] == 201 assert result["status_code"] == 201
@@ -644,6 +714,7 @@ async def test_service_execute_action_companies_post(db_session):
async def test_service_execute_action_companies_patch(db_session): async def test_service_execute_action_companies_patch(db_session):
"""Service: execute_action with PATCH /api/v1/companies/{id} updates company.""" """Service: execute_action with PATCH /api/v1/companies/{id} updates company."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -653,15 +724,27 @@ async def test_service_execute_action_companies_patch(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
create_result = await execute_action( create_result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "POST", "path": "/api/v1/companies", "body": {"name": "PatchCo"}}, {"method": "POST", "path": "/api/v1/companies", "body": {"name": "PatchCo"}},
) )
company_id = create_result["data"]["id"] company_id = create_result["data"]["id"]
# Now patch it # Now patch it
patch_result = await execute_action( patch_result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
{"method": "PATCH", "path": f"/api/v1/companies/{company_id}", "body": {"name": "PatchedCo"}}, tenant_id,
admin_id,
"admin",
conv_id,
{
"method": "PATCH",
"path": f"/api/v1/companies/{company_id}",
"body": {"name": "PatchedCo"},
},
) )
assert patch_result["success"] is True assert patch_result["success"] is True
assert patch_result["data"]["name"] == "PatchedCo" assert patch_result["data"]["name"] == "PatchedCo"
@@ -671,6 +754,7 @@ async def test_service_execute_action_companies_patch(db_session):
async def test_service_execute_action_companies_patch_not_found(db_session): async def test_service_execute_action_companies_patch_not_found(db_session):
"""Service: execute_action with PATCH non-existent company returns 404.""" """Service: execute_action with PATCH non-existent company returns 404."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -679,8 +763,16 @@ async def test_service_execute_action_companies_patch_not_found(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
{"method": "PATCH", "path": "/api/v1/companies/00000000-0000-0000-0000-000000000000", "body": {"name": "X"}}, tenant_id,
admin_id,
"admin",
conv_id,
{
"method": "PATCH",
"path": "/api/v1/companies/00000000-0000-0000-0000-000000000000",
"body": {"name": "X"},
},
) )
assert result["success"] is False assert result["success"] is False
assert result["status_code"] == 404 assert result["status_code"] == 404
@@ -690,6 +782,7 @@ async def test_service_execute_action_companies_patch_not_found(db_session):
async def test_service_execute_action_companies_patch_no_id(db_session): async def test_service_execute_action_companies_patch_no_id(db_session):
"""Service: execute_action with PATCH /api/v1/companies/{id} returns 400.""" """Service: execute_action with PATCH /api/v1/companies/{id} returns 400."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -698,7 +791,11 @@ async def test_service_execute_action_companies_patch_no_id(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "PATCH", "path": "/api/v1/companies/{id}", "body": {"name": "X"}}, {"method": "PATCH", "path": "/api/v1/companies/{id}", "body": {"name": "X"}},
) )
assert result["success"] is False assert result["success"] is False
@@ -709,6 +806,7 @@ async def test_service_execute_action_companies_patch_no_id(db_session):
async def test_service_execute_action_companies_delete(db_session): async def test_service_execute_action_companies_delete(db_session):
"""Service: execute_action with DELETE /api/v1/companies/{id} soft-deletes company.""" """Service: execute_action with DELETE /api/v1/companies/{id} soft-deletes company."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -717,13 +815,21 @@ async def test_service_execute_action_companies_delete(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
create_result = await execute_action( create_result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "POST", "path": "/api/v1/companies", "body": {"name": "DeleteMe"}}, {"method": "POST", "path": "/api/v1/companies", "body": {"name": "DeleteMe"}},
) )
company_id = create_result["data"]["id"] company_id = create_result["data"]["id"]
del_result = await execute_action( del_result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "DELETE", "path": f"/api/v1/companies/{company_id}", "body": None}, {"method": "DELETE", "path": f"/api/v1/companies/{company_id}", "body": None},
) )
assert del_result["success"] is True assert del_result["success"] is True
@@ -734,6 +840,7 @@ async def test_service_execute_action_companies_delete(db_session):
async def test_service_execute_action_companies_delete_not_found(db_session): async def test_service_execute_action_companies_delete_not_found(db_session):
"""Service: execute_action with DELETE non-existent company returns 404.""" """Service: execute_action with DELETE non-existent company returns 404."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -742,8 +849,16 @@ async def test_service_execute_action_companies_delete_not_found(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
{"method": "DELETE", "path": "/api/v1/companies/00000000-0000-0000-0000-000000000000", "body": None}, tenant_id,
admin_id,
"admin",
conv_id,
{
"method": "DELETE",
"path": "/api/v1/companies/00000000-0000-0000-0000-000000000000",
"body": None,
},
) )
assert result["success"] is False assert result["success"] is False
assert result["status_code"] == 404 assert result["status_code"] == 404
@@ -753,6 +868,7 @@ async def test_service_execute_action_companies_delete_not_found(db_session):
async def test_service_execute_action_companies_delete_no_id(db_session): async def test_service_execute_action_companies_delete_no_id(db_session):
"""Service: execute_action with DELETE without ID returns 400.""" """Service: execute_action with DELETE without ID returns 400."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -761,7 +877,11 @@ async def test_service_execute_action_companies_delete_no_id(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "DELETE", "path": "/api/v1/companies/{id}", "body": None}, {"method": "DELETE", "path": "/api/v1/companies/{id}", "body": None},
) )
assert result["success"] is False assert result["success"] is False
@@ -772,6 +892,7 @@ async def test_service_execute_action_companies_delete_no_id(db_session):
async def test_service_execute_action_contacts_get(db_session): async def test_service_execute_action_contacts_get(db_session):
"""Service: execute_action with GET /api/v1/contacts returns list.""" """Service: execute_action with GET /api/v1/contacts returns list."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -780,7 +901,11 @@ async def test_service_execute_action_contacts_get(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "GET", "path": "/api/v1/contacts", "body": None}, {"method": "GET", "path": "/api/v1/contacts", "body": None},
) )
assert result["success"] is True assert result["success"] is True
@@ -792,6 +917,7 @@ async def test_service_execute_action_contacts_post(db_session):
"""Service: execute_action with POST /api/v1/contacts — Contact model uses first_name/last_name, """Service: execute_action with POST /api/v1/contacts — Contact model uses first_name/last_name,
service passes 'name' which causes error, verify error handling works.""" service passes 'name' which causes error, verify error handling works."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -800,8 +926,16 @@ async def test_service_execute_action_contacts_post(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
{"method": "POST", "path": "/api/v1/contacts", "body": {"name": "John Doe", "email": "john@example.com"}}, tenant_id,
admin_id,
"admin",
conv_id,
{
"method": "POST",
"path": "/api/v1/contacts",
"body": {"name": "John Doe", "email": "john@example.com"},
},
) )
# Contact model has first_name/last_name, not name — service code attempts to set 'name' # Contact model has first_name/last_name, not name — service code attempts to set 'name'
# which raises TypeError, caught by execute_action's try/except # which raises TypeError, caught by execute_action's try/except
@@ -813,6 +947,7 @@ async def test_service_execute_action_contacts_post(db_session):
async def test_service_execute_action_contacts_unsupported_method(db_session): async def test_service_execute_action_contacts_unsupported_method(db_session):
"""Service: execute_action with DELETE /api/v1/contacts returns 400 (unsupported).""" """Service: execute_action with DELETE /api/v1/contacts returns 400 (unsupported)."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -821,7 +956,11 @@ async def test_service_execute_action_contacts_unsupported_method(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "DELETE", "path": "/api/v1/contacts/123", "body": None}, {"method": "DELETE", "path": "/api/v1/contacts/123", "body": None},
) )
assert result["success"] is False assert result["success"] is False
@@ -832,6 +971,7 @@ async def test_service_execute_action_contacts_unsupported_method(db_session):
async def test_service_execute_action_workflows_get(db_session): async def test_service_execute_action_workflows_get(db_session):
"""Service: execute_action with GET /api/v1/workflows returns list.""" """Service: execute_action with GET /api/v1/workflows returns list."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -840,7 +980,11 @@ async def test_service_execute_action_workflows_get(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "GET", "path": "/api/v1/workflows", "body": None}, {"method": "GET", "path": "/api/v1/workflows", "body": None},
) )
assert result["success"] is True assert result["success"] is True
@@ -851,6 +995,7 @@ async def test_service_execute_action_workflows_get(db_session):
async def test_service_execute_action_workflows_unsupported_method(db_session): async def test_service_execute_action_workflows_unsupported_method(db_session):
"""Service: execute_action with POST /api/v1/workflows returns 400 (unsupported).""" """Service: execute_action with POST /api/v1/workflows returns 400 (unsupported)."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -859,7 +1004,11 @@ async def test_service_execute_action_workflows_unsupported_method(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "POST", "path": "/api/v1/workflows", "body": {"name": "test"}}, {"method": "POST", "path": "/api/v1/workflows", "body": {"name": "test"}},
) )
assert result["success"] is False assert result["success"] is False
@@ -870,6 +1019,7 @@ async def test_service_execute_action_workflows_unsupported_method(db_session):
async def test_service_execute_action_unsupported_entity(db_session): async def test_service_execute_action_unsupported_entity(db_session):
"""Service: execute_action with unknown entity returns 400.""" """Service: execute_action with unknown entity returns 400."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -878,7 +1028,11 @@ async def test_service_execute_action_unsupported_entity(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "GET", "path": "/api/v1/unknown", "body": None}, {"method": "GET", "path": "/api/v1/unknown", "body": None},
) )
assert result["success"] is False assert result["success"] is False
@@ -890,6 +1044,7 @@ async def test_service_execute_action_unsupported_entity(db_session):
async def test_service_execute_action_companies_unsupported_method(db_session): async def test_service_execute_action_companies_unsupported_method(db_session):
"""Service: execute_action with PUT /api/v1/companies returns 400 (unsupported).""" """Service: execute_action with PUT /api/v1/companies returns 400 (unsupported)."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -898,7 +1053,11 @@ async def test_service_execute_action_companies_unsupported_method(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", conv_id, db_session,
tenant_id,
admin_id,
"admin",
conv_id,
{"method": "PUT", "path": "/api/v1/companies", "body": {}}, {"method": "PUT", "path": "/api/v1/companies", "body": {}},
) )
assert result["success"] is False assert result["success"] is False
@@ -909,12 +1068,16 @@ async def test_service_execute_action_companies_unsupported_method(db_session):
async def test_service_execute_action_invalid_conversation(db_session): async def test_service_execute_action_invalid_conversation(db_session):
"""Service: execute_action with invalid conversation_id returns 404.""" """Service: execute_action with invalid conversation_id returns 404."""
from app.services.ai_copilot_service import execute_action from app.services.ai_copilot_service import execute_action
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "admin", db_session,
tenant_id,
admin_id,
"admin",
"00000000-0000-0000-0000-000000000000", "00000000-0000-0000-0000-000000000000",
{"method": "GET", "path": "/api/v1/companies", "body": None}, {"method": "GET", "path": "/api/v1/companies", "body": None},
) )
@@ -926,6 +1089,7 @@ async def test_service_execute_action_invalid_conversation(db_session):
async def test_service_execute_action_rbac_blocked(db_session): async def test_service_execute_action_rbac_blocked(db_session):
"""Service: execute_action as viewer with DELETE returns 403.""" """Service: execute_action as viewer with DELETE returns 403."""
from app.services.ai_copilot_service import execute_action, process_query from app.services.ai_copilot_service import execute_action, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -934,8 +1098,16 @@ async def test_service_execute_action_rbac_blocked(db_session):
conv_id = query_result["conversation_id"] conv_id = query_result["conversation_id"]
result = await execute_action( result = await execute_action(
db_session, tenant_id, admin_id, "viewer", conv_id, db_session,
{"method": "DELETE", "path": "/api/v1/companies/00000000-0000-0000-0000-000000000000", "body": None}, tenant_id,
admin_id,
"viewer",
conv_id,
{
"method": "DELETE",
"path": "/api/v1/companies/00000000-0000-0000-0000-000000000000",
"body": None,
},
) )
assert result["status_code"] == 403 assert result["status_code"] == 403
assert result["success"] is False assert result["success"] is False
@@ -945,6 +1117,7 @@ async def test_service_execute_action_rbac_blocked(db_session):
async def test_service_get_history_pagination(db_session): async def test_service_get_history_pagination(db_session):
"""Service: get_history returns paginated results.""" """Service: get_history returns paginated results."""
from app.services.ai_copilot_service import get_history, process_query from app.services.ai_copilot_service import get_history, process_query
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -963,6 +1136,7 @@ async def test_service_get_history_pagination(db_session):
async def test_service_get_history_empty(db_session): async def test_service_get_history_empty(db_session):
"""Service: get_history returns empty when no conversations exist.""" """Service: get_history returns empty when no conversations exist."""
from app.services.ai_copilot_service import get_history from app.services.ai_copilot_service import get_history
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
tenant_id = seed["tenant_a"].id tenant_id = seed["tenant_a"].id
admin_id = seed["admin_a"].id admin_id = seed["admin_a"].id
@@ -975,6 +1149,7 @@ async def test_service_get_history_empty(db_session):
def test_service_derive_rbac_from_path_companies_get(): def test_service_derive_rbac_from_path_companies_get():
"""Service: _derive_rbac_from_path for GET /api/v1/companies → (companies, read).""" """Service: _derive_rbac_from_path for GET /api/v1/companies → (companies, read)."""
from app.services.ai_copilot_service import _derive_rbac_from_path from app.services.ai_copilot_service import _derive_rbac_from_path
module, action = _derive_rbac_from_path("GET", "/api/v1/companies") module, action = _derive_rbac_from_path("GET", "/api/v1/companies")
assert module == "companies" assert module == "companies"
assert action == "read" assert action == "read"
@@ -983,6 +1158,7 @@ def test_service_derive_rbac_from_path_companies_get():
def test_service_derive_rbac_from_path_contacts_post(): def test_service_derive_rbac_from_path_contacts_post():
"""Service: _derive_rbac_from_path for POST /api/v1/contacts → (contacts, create).""" """Service: _derive_rbac_from_path for POST /api/v1/contacts → (contacts, create)."""
from app.services.ai_copilot_service import _derive_rbac_from_path from app.services.ai_copilot_service import _derive_rbac_from_path
module, action = _derive_rbac_from_path("POST", "/api/v1/contacts") module, action = _derive_rbac_from_path("POST", "/api/v1/contacts")
assert module == "contacts" assert module == "contacts"
assert action == "create" assert action == "create"
@@ -991,6 +1167,7 @@ def test_service_derive_rbac_from_path_contacts_post():
def test_service_derive_rbac_from_path_workflows_patch(): def test_service_derive_rbac_from_path_workflows_patch():
"""Service: _derive_rbac_from_path for PATCH /api/v1/workflows → (workflows, update).""" """Service: _derive_rbac_from_path for PATCH /api/v1/workflows → (workflows, update)."""
from app.services.ai_copilot_service import _derive_rbac_from_path from app.services.ai_copilot_service import _derive_rbac_from_path
module, action = _derive_rbac_from_path("PATCH", "/api/v1/workflows") module, action = _derive_rbac_from_path("PATCH", "/api/v1/workflows")
assert module == "workflows" assert module == "workflows"
assert action == "update" assert action == "update"
@@ -999,6 +1176,7 @@ def test_service_derive_rbac_from_path_workflows_patch():
def test_service_derive_rbac_from_path_unknown_entity(): def test_service_derive_rbac_from_path_unknown_entity():
"""Service: _derive_rbac_from_path for unknown entity returns entity as module.""" """Service: _derive_rbac_from_path for unknown entity returns entity as module."""
from app.services.ai_copilot_service import _derive_rbac_from_path from app.services.ai_copilot_service import _derive_rbac_from_path
module, action = _derive_rbac_from_path("DELETE", "/api/v1/foobar") module, action = _derive_rbac_from_path("DELETE", "/api/v1/foobar")
assert module == "foobar" assert module == "foobar"
assert action == "delete" assert action == "delete"
@@ -1007,6 +1185,7 @@ def test_service_derive_rbac_from_path_unknown_entity():
def test_service_derive_rbac_from_path_ai_entity(): def test_service_derive_rbac_from_path_ai_entity():
"""Service: _derive_rbac_from_path for /api/v1/ai → (ai_copilot, read).""" """Service: _derive_rbac_from_path for /api/v1/ai → (ai_copilot, read)."""
from app.services.ai_copilot_service import _derive_rbac_from_path from app.services.ai_copilot_service import _derive_rbac_from_path
module, action = _derive_rbac_from_path("GET", "/api/v1/ai/copilot/query") module, action = _derive_rbac_from_path("GET", "/api/v1/ai/copilot/query")
assert module == "ai_copilot" assert module == "ai_copilot"
assert action == "read" assert action == "read"
@@ -1015,14 +1194,17 @@ def test_service_derive_rbac_from_path_ai_entity():
def test_service_safe_iso_none(): def test_service_safe_iso_none():
"""Service: _safe_iso with None returns None.""" """Service: _safe_iso with None returns None."""
from app.services.ai_copilot_service import _safe_iso from app.services.ai_copilot_service import _safe_iso
assert _safe_iso(None) is None assert _safe_iso(None) is None
def test_service_safe_iso_datetime(): def test_service_safe_iso_datetime():
"""Service: _safe_iso with datetime returns ISO string.""" """Service: _safe_iso with datetime returns ISO string."""
from datetime import datetime
from app.services.ai_copilot_service import _safe_iso from app.services.ai_copilot_service import _safe_iso
from datetime import datetime, timezone
dt = datetime(2026, 1, 1, 12, 0, 0, tzinfo=timezone.utc) dt = datetime(2026, 1, 1, 12, 0, 0, tzinfo=UTC)
result = _safe_iso(dt) result = _safe_iso(dt)
assert result is not None assert result is not None
assert "2026-01-01" in result assert "2026-01-01" in result
@@ -1031,15 +1213,18 @@ def test_service_safe_iso_datetime():
def test_service_safe_iso_exception(): def test_service_safe_iso_exception():
"""Service: _safe_iso with object that raises on isoformat returns None.""" """Service: _safe_iso with object that raises on isoformat returns None."""
from app.services.ai_copilot_service import _safe_iso from app.services.ai_copilot_service import _safe_iso
class Bad: class Bad:
def isoformat(self): def isoformat(self):
raise ValueError("bad") raise ValueError("bad")
assert _safe_iso(Bad()) is None assert _safe_iso(Bad()) is None
def test_service_get_attr_missing(): def test_service_get_attr_missing():
"""Service: _get_attr with missing attribute returns default.""" """Service: _get_attr with missing attribute returns default."""
from app.services.ai_copilot_service import _get_attr from app.services.ai_copilot_service import _get_attr
obj = type("Obj", (), {"x": 1})() obj = type("Obj", (), {"x": 1})()
assert _get_attr(obj, "x") == 1 assert _get_attr(obj, "x") == 1
assert _get_attr(obj, "y", "default") == "default" assert _get_attr(obj, "y", "default") == "default"
@@ -1048,12 +1233,14 @@ def test_service_get_attr_missing():
def test_service_get_attr_none(): def test_service_get_attr_none():
"""Service: _get_attr with None value returns default.""" """Service: _get_attr with None value returns default."""
from app.services.ai_copilot_service import _get_attr from app.services.ai_copilot_service import _get_attr
obj = type("Obj", (), {"x": None})() obj = type("Obj", (), {"x": None})()
assert _get_attr(obj, "x", "fallback") == "fallback" assert _get_attr(obj, "x", "fallback") == "fallback"
# ─── Route Error Path Tests for AI Copilot ─── # ─── Route Error Path Tests for AI Copilot ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_route_copilot_execute_not_found(client: AsyncClient, db_session): async def test_route_copilot_execute_not_found(client: AsyncClient, db_session):
"""Route: POST /execute with invalid conversation_id returns 404.""" """Route: POST /execute with invalid conversation_id returns 404."""
@@ -1064,7 +1251,12 @@ async def test_route_copilot_execute_not_found(client: AsyncClient, db_session):
"/api/v1/ai/copilot/execute", "/api/v1/ai/copilot/execute",
json={ json={
"conversation_id": "00000000-0000-0000-0000-000000000000", "conversation_id": "00000000-0000-0000-0000-000000000000",
"action": {"method": "GET", "path": "/api/v1/companies", "description": "List", "confidence": 0.9}, "action": {
"method": "GET",
"path": "/api/v1/companies",
"description": "List",
"confidence": 0.9,
},
}, },
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
@@ -1080,7 +1272,10 @@ async def test_route_copilot_execute_rbac_blocked(client: AsyncClient, db_sessio
# First create a conversation as viewer # First create a conversation as viewer
query_resp = await client.post( query_resp = await client.post(
"/api/v1/ai/copilot/query", "/api/v1/ai/copilot/query",
json={"query": "Delete company", "context": {"entity_id": "00000000-0000-0000-0000-000000000000"}}, json={
"query": "Delete company",
"context": {"entity_id": "00000000-0000-0000-0000-000000000000"},
},
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
conv_id = query_resp.json()["conversation_id"] conv_id = query_resp.json()["conversation_id"]
@@ -1108,7 +1303,12 @@ async def test_route_copilot_execute_unauthenticated(client: AsyncClient, db_ses
"/api/v1/ai/copilot/execute", "/api/v1/ai/copilot/execute",
json={ json={
"conversation_id": "00000000-0000-0000-0000-000000000000", "conversation_id": "00000000-0000-0000-0000-000000000000",
"action": {"method": "GET", "path": "/api/v1/companies", "description": "List", "confidence": 0.9}, "action": {
"method": "GET",
"path": "/api/v1/companies",
"description": "List",
"confidence": 0.9,
},
}, },
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
+12 -5
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -21,8 +21,9 @@ class TestAuthLogin:
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
assert resp.status_code == 200 assert resp.status_code == 200
assert "leocrm_session" in resp.headers.get("set-cookie", "").lower() or \ assert "leocrm_session" in resp.headers.get(
"leocrm_session" in str(resp.cookies) "set-cookie", ""
).lower() or "leocrm_session" in str(resp.cookies)
data = resp.json() data = resp.json()
assert data["email"] == "admin@tenanta.com" assert data["email"] == "admin@tenanta.com"
assert data["role"] == "admin" assert data["role"] == "admin"
@@ -64,7 +65,9 @@ class TestAuthMe:
class TestAuthLogout: class TestAuthLogout:
"""AC 5: logout.""" """AC 5: logout."""
async def test_logout_returns_200_and_invalidates_session(self, client: AsyncClient, db_session): async def test_logout_returns_200_and_invalidates_session(
self, client: AsyncClient, db_session
):
"""AC 5: POST /api/v1/auth/logout -> 200, session invalidated.""" """AC 5: POST /api/v1/auth/logout -> 200, session invalidated."""
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
@@ -104,6 +107,7 @@ class TestPasswordReset:
async def test_password_reset_confirm_valid_token(self, client: AsyncClient, db_session): async def test_password_reset_confirm_valid_token(self, client: AsyncClient, db_session):
"""AC 7: POST /api/v1/auth/password-reset/confirm valid token -> 200.""" """AC 7: POST /api/v1/auth/password-reset/confirm valid token -> 200."""
from app.services.auth_service import auth_service from app.services.auth_service import auth_service
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
raw_token = await auth_service.get_password_reset_token_raw(db_session, "admin@tenanta.com") raw_token = await auth_service.get_password_reset_token_raw(db_session, "admin@tenanta.com")
await db_session.commit() await db_session.commit()
@@ -119,6 +123,7 @@ class TestPasswordReset:
async def test_password_reset_confirm_expired_token(self, client: AsyncClient, db_session): async def test_password_reset_confirm_expired_token(self, client: AsyncClient, db_session):
"""AC 8: POST /api/v1/auth/password-reset/confirm expired token -> 400.""" """AC 8: POST /api/v1/auth/password-reset/confirm expired token -> 400."""
from app.services.auth_service import auth_service from app.services.auth_service import auth_service
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
raw_token = await auth_service.create_expired_reset_token(db_session, "admin@tenanta.com") raw_token = await auth_service.create_expired_reset_token(db_session, "admin@tenanta.com")
await db_session.commit() await db_session.commit()
@@ -147,8 +152,10 @@ class TestSwitchTenant:
initial_tenant = resp.json()["tenant_id"] initial_tenant = resp.json()["tenant_id"]
# Get tenant B ID # Get tenant B ID
from app.models.tenant import Tenant
from sqlalchemy import select from sqlalchemy import select
from app.models.tenant import Tenant
q = select(Tenant).where(Tenant.slug == "tenant-b") q = select(Tenant).where(Tenant.slug == "tenant-b")
result = await db_session.execute(q) result = await db_session.execute(q)
tenant_b = result.scalar_one() tenant_b = result.scalar_one()
+4 -2
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -186,7 +186,7 @@ class TestCompanyDelete:
json={"first_name": "John", "last_name": "Doe", "company_ids": [company_id]}, json={"first_name": "John", "last_name": "Doe", "company_ids": [company_id]},
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
contact_id = cont_resp.json()["id"] cont_resp.json()["id"]
# Verify link exists # Verify link exists
detail_resp = await client.get(f"/api/v1/companies/{company_id}", headers=ORIGIN_HEADER) detail_resp = await client.get(f"/api/v1/companies/{company_id}", headers=ORIGIN_HEADER)
assert len(detail_resp.json()["contacts"]) == 1 assert len(detail_resp.json()["contacts"]) == 1
@@ -318,7 +318,9 @@ class TestCompanyAuditAndSoftDelete:
async def test_audit_log_on_company_create(self, client: AsyncClient, db_session): async def test_audit_log_on_company_create(self, client: AsyncClient, db_session):
"""AC 23: Audit log entry on every company/contact mutation.""" """AC 23: Audit log entry on every company/contact mutation."""
from sqlalchemy import select from sqlalchemy import select
from app.models.audit import AuditLog from app.models.audit import AuditLog
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
await client.post( await client.post(
+11 -4
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -29,7 +29,9 @@ class TestContactList:
class TestContactCreate: class TestContactCreate:
"""AC 15: Create contact with company_ids array -> N:M links.""" """AC 15: Create contact with company_ids array -> N:M links."""
async def test_create_contact_with_company_ids_returns_201(self, client: AsyncClient, db_session): async def test_create_contact_with_company_ids_returns_201(
self, client: AsyncClient, db_session
):
"""AC 15: POST /api/v1/contacts mit company_ids array -> 201 + N:M links.""" """AC 15: POST /api/v1/contacts mit company_ids array -> 201 + N:M links."""
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
@@ -132,12 +134,17 @@ class TestContactDelete:
names = [f"{item['first_name']} {item['last_name']}" for item in list_resp.json()["items"]] names = [f"{item['first_name']} {item['last_name']}" for item in list_resp.json()["items"]]
assert "Delete Me" not in names assert "Delete Me" not in names
async def test_delete_contact_gdpr_hard_delete_returns_204(self, client: AsyncClient, db_session): async def test_delete_contact_gdpr_hard_delete_returns_204(
self, client: AsyncClient, db_session
):
"""AC 19: DELETE /api/v1/contacts/{id}?gdpr=true -> 204, hard-delete + deletion_log.""" """AC 19: DELETE /api/v1/contacts/{id}?gdpr=true -> 204, hard-delete + deletion_log."""
import uuid as uuid_mod
from sqlalchemy import select from sqlalchemy import select
from app.models.audit import DeletionLog from app.models.audit import DeletionLog
from app.models.contact import Contact from app.models.contact import Contact
import uuid as uuid_mod
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
create_resp = await client.post( create_resp = await client.post(
+22 -14
View File
@@ -7,16 +7,16 @@ import uuid
import pytest import pytest
import pytest_asyncio import pytest_asyncio
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, AsyncSession from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
from app.core.db import reset_engine_for_testing, close_engine from app.core.db import close_engine, reset_engine_for_testing
from app.core.event_bus import get_event_bus from app.core.event_bus import get_event_bus
from app.core.service_container import get_container from app.core.service_container import get_container
from app.main import create_app from app.main import create_app
from app.plugins.registry import reset_registry_for_testing
from app.plugins.builtins.entity_links import EntityLinksPlugin from app.plugins.builtins.entity_links import EntityLinksPlugin
from app.plugins.registry import reset_registry_for_testing
from app.services.plugin_service import reset_plugin_service_for_testing from app.services.plugin_service import reset_plugin_service_for_testing
from tests.conftest import seed_tenant_and_users, login_client, ORIGIN_HEADER from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest_asyncio.fixture @pytest_asyncio.fixture
@@ -240,14 +240,18 @@ async def test_event_cleanup_on_company_deleted(authed_client: AsyncClient):
# Publish company.deleted event # Publish company.deleted event
event_bus = get_event_bus() event_bus = get_event_bus()
await event_bus.publish("company.deleted", { await event_bus.publish(
"entity_id": str(company_id), "company.deleted",
"company_id": str(company_id), {
"tenant_id": str(tenant_id), "entity_id": str(company_id),
}) "company_id": str(company_id),
"tenant_id": str(tenant_id),
},
)
# Allow async event handler to complete # Allow async event handler to complete
import asyncio import asyncio
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
# Verify link is cleaned up # Verify link is cleaned up
@@ -279,14 +283,18 @@ async def test_event_cleanup_on_contact_deleted(authed_client: AsyncClient):
# Publish contact.deleted event # Publish contact.deleted event
event_bus = get_event_bus() event_bus = get_event_bus()
await event_bus.publish("contact.deleted", { await event_bus.publish(
"entity_id": str(contact_id), "contact.deleted",
"contact_id": str(contact_id), {
"tenant_id": str(tenant_id), "entity_id": str(contact_id),
}) "contact_id": str(contact_id),
"tenant_id": str(tenant_id),
},
)
# Allow async event handler to complete # Allow async event handler to complete
import asyncio import asyncio
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
# Verify link is cleaned up # Verify link is cleaned up
+1 -3
View File
@@ -2,12 +2,10 @@
from __future__ import annotations from __future__ import annotations
import io
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
CSV_COMPANIES = """name,industry,phone,email,website,description CSV_COMPANIES = """name,industry,phone,email,website,description
ImportCorp,IT,123456,import@example.com,https://import.example,Imported company ImportCorp,IT,123456,import@example.com,https://import.example,Imported company
+35 -16
View File
@@ -2,16 +2,14 @@
from __future__ import annotations from __future__ import annotations
import uuid from datetime import UTC, datetime
from datetime import datetime, timezone
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from sqlalchemy.ext.asyncio import AsyncSession
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client
from app.core.notifications import create_notification from app.core.notifications import create_notification
from app.models.notification import Notification from app.models.notification import Notification
from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -23,22 +21,31 @@ class TestNotifications:
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
# Create some notifications for admin_a # Create some notifications for admin_a
await create_notification( await create_notification(
db_session, seed["tenant_a"].id, seed["admin_a"].id, db_session,
"info", "Read Notif", "Already read", seed["tenant_a"].id,
seed["admin_a"].id,
"info",
"Read Notif",
"Already read",
) )
await db_session.flush() await db_session.flush()
# Mark the first one as read # Mark the first one as read
from sqlalchemy import select, update from sqlalchemy import select
q = select(Notification).where(Notification.title == "Read Notif") q = select(Notification).where(Notification.title == "Read Notif")
result = await db_session.execute(q) result = await db_session.execute(q)
first_notif = result.scalar_one() first_notif = result.scalar_one()
first_notif.read_at = datetime.now(timezone.utc) first_notif.read_at = datetime.now(UTC)
await db_session.flush() await db_session.flush()
# Create an unread one # Create an unread one
await create_notification( await create_notification(
db_session, seed["tenant_a"].id, seed["admin_a"].id, db_session,
"info", "Unread Notif", "Not read yet", seed["tenant_a"].id,
seed["admin_a"].id,
"info",
"Unread Notif",
"Not read yet",
) )
await db_session.commit() await db_session.commit()
@@ -64,8 +71,12 @@ class TestNotifications:
"""AC 25: PATCH /api/v1/notifications/{id}/read -> 200.""" """AC 25: PATCH /api/v1/notifications/{id}/read -> 200."""
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
notif = await create_notification( notif = await create_notification(
db_session, seed["tenant_a"].id, seed["admin_a"].id, db_session,
"info", "Test Notif", "Test body", seed["tenant_a"].id,
seed["admin_a"].id,
"info",
"Test Notif",
"Test body",
) )
await db_session.commit() await db_session.commit()
@@ -81,12 +92,20 @@ class TestNotifications:
"""AC 26: GET /api/v1/notifications/unread-count -> 200 + integer count.""" """AC 26: GET /api/v1/notifications/unread-count -> 200 + integer count."""
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
await create_notification( await create_notification(
db_session, seed["tenant_a"].id, seed["admin_a"].id, db_session,
"info", "Unread 1", "Body 1", seed["tenant_a"].id,
seed["admin_a"].id,
"info",
"Unread 1",
"Body 1",
) )
await create_notification( await create_notification(
db_session, seed["tenant_a"].id, seed["admin_a"].id, db_session,
"info", "Unread 2", "Body 2", seed["tenant_a"].id,
seed["admin_a"].id,
"info",
"Unread 2",
"Body 2",
) )
await db_session.commit() await db_session.commit()
+13 -10
View File
@@ -3,20 +3,20 @@
from __future__ import annotations from __future__ import annotations
import uuid import uuid
from datetime import datetime, timedelta, timezone from datetime import UTC, datetime, timedelta
import pytest import pytest
import pytest_asyncio import pytest_asyncio
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, AsyncSession from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
from app.core.db import reset_engine_for_testing, close_engine from app.core.db import close_engine, reset_engine_for_testing
from app.core.service_container import get_container from app.core.service_container import get_container
from app.main import create_app from app.main import create_app
from app.plugins.registry import reset_registry_for_testing
from app.plugins.builtins.permissions import PermissionsPlugin from app.plugins.builtins.permissions import PermissionsPlugin
from app.plugins.registry import reset_registry_for_testing
from app.services.plugin_service import reset_plugin_service_for_testing from app.services.plugin_service import reset_plugin_service_for_testing
from tests.conftest import seed_tenant_and_users, login_client, ORIGIN_HEADER from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest_asyncio.fixture @pytest_asyncio.fixture
@@ -182,7 +182,7 @@ async def test_expired_share_link(authed_client: AsyncClient):
file_id = str(uuid.uuid4()) file_id = str(uuid.uuid4())
# Create link with expiry in the past # Create link with expiry in the past
past = datetime.now(timezone.utc) - timedelta(hours=1) past = datetime.now(UTC) - timedelta(hours=1)
resp = await client.post( resp = await client.post(
f"/api/v1/dms/files/{file_id}/share-link", f"/api/v1/dms/files/{file_id}/share-link",
json={"expires_at": past.isoformat(), "access_level": "download"}, json={"expires_at": past.isoformat(), "access_level": "download"},
@@ -235,21 +235,23 @@ async def test_revoke_share_link(authed_client: AsyncClient):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_permission_403_for_unauthorized_user(plugin_client: AsyncClient, db_session: AsyncSession): async def test_permission_403_for_unauthorized_user(
plugin_client: AsyncClient, db_session: AsyncSession
):
"""AC14: Folder permissions enforced: user without read → 403. """AC14: Folder permissions enforced: user without read → 403.
The permissions plugin does not intercept file access (no DMS file system yet). The permissions plugin does not intercept file access (no DMS file system yet).
We test the permission check helper directly: a user with no permission record We test the permission check helper directly: a user with no permission record
gets denied (False), which translates to 403 at the route layer. gets denied (False), which translates to 403 at the route layer.
""" """
from app.plugins.builtins.permissions.routes import check_user_file_permission
from app.core.db import get_session_factory from app.core.db import get_session_factory
from app.plugins.builtins.permissions.routes import check_user_file_permission
seed = await seed_tenant_and_users(db_session) seed = await seed_tenant_and_users(db_session)
await login_client(plugin_client, "admin@tenanta.com") await login_client(plugin_client, "admin@tenanta.com")
# Install + activate permissions plugin # Install + activate permissions plugin
registry = reset_registry_for_testing() reset_registry_for_testing()
# We need to set up the registry properly # We need to set up the registry properly
# Since plugin_client fixture wasn't used, manually install # Since plugin_client fixture wasn't used, manually install
resp = await plugin_client.post("/api/v1/plugins/permissions/install", headers=ORIGIN_HEADER) resp = await plugin_client.post("/api/v1/plugins/permissions/install", headers=ORIGIN_HEADER)
@@ -270,6 +272,7 @@ async def test_permission_403_for_unauthorized_user(plugin_client: AsyncClient,
# Grant read to admin # Grant read to admin
from app.plugins.builtins.permissions.models import Permission from app.plugins.builtins.permissions.models import Permission
perm = Permission( perm = Permission(
tenant_id=tenant_id, tenant_id=tenant_id,
file_id=file_id, file_id=file_id,
@@ -422,7 +425,7 @@ async def test_public_access_post_expired(authed_client: AsyncClient):
"""POST /api/public/share/{token} with expired link → 410.""" """POST /api/public/share/{token} with expired link → 410."""
client, seed = authed_client client, seed = authed_client
file_id = str(uuid.uuid4()) file_id = str(uuid.uuid4())
past = datetime.now(timezone.utc) - timedelta(hours=1) past = datetime.now(UTC) - timedelta(hours=1)
resp = await client.post( resp = await client.post(
f"/api/v1/dms/files/{file_id}/share-link", f"/api/v1/dms/files/{file_id}/share-link",
json={"expires_at": past.isoformat(), "access_level": "download"}, json={"expires_at": past.isoformat(), "access_level": "download"},
+163 -52
View File
@@ -3,31 +3,31 @@
from __future__ import annotations from __future__ import annotations
import asyncio import asyncio
import uuid
from typing import Any from typing import Any
import pytest import pytest
import pytest_asyncio import pytest_asyncio
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy import text, select, inspect from sqlalchemy import inspect, select, text
from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, AsyncSession from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker
from app.core.db import Base, reset_engine_for_testing, close_engine, get_engine from app.core.db import close_engine, reset_engine_for_testing
from app.core.event_bus import get_event_bus from app.core.event_bus import get_event_bus
from app.core.service_container import get_container from app.core.service_container import get_container
from app.main import create_app from app.main import create_app
from app.models.plugin import Plugin as PluginModel, PluginMigration from app.models.plugin import Plugin as PluginModel
from app.models.plugin import PluginMigration
from app.plugins.base import BasePlugin from app.plugins.base import BasePlugin
from app.plugins.manifest import PluginManifest
from app.plugins.registry import get_registry, reset_registry_for_testing
from app.plugins.migration_runner import MigrationRunner, MigrationValidationError
from app.plugins.builtins.test_sample import TestSamplePlugin from app.plugins.builtins.test_sample import TestSamplePlugin
from app.plugins.manifest import PluginManifest
from app.plugins.migration_runner import MigrationRunner, MigrationValidationError
from app.plugins.registry import get_registry, reset_registry_for_testing
from app.services.plugin_service import reset_plugin_service_for_testing from app.services.plugin_service import reset_plugin_service_for_testing
from tests.conftest import TEST_DB_URL, seed_tenant_and_users, login_client, ORIGIN_HEADER from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
# ─── Bad Plugin for AC11 (migration without tenant_id) ─── # ─── Bad Plugin for AC11 (migration without tenant_id) ───
class BadMigrationPlugin(BasePlugin): class BadMigrationPlugin(BasePlugin):
"""Plugin with a migration that creates a table WITHOUT tenant_id.""" """Plugin with a migration that creates a table WITHOUT tenant_id."""
@@ -46,6 +46,7 @@ class BadMigrationPlugin(BasePlugin):
# ─── Fixtures ─── # ─── Fixtures ───
@pytest_asyncio.fixture @pytest_asyncio.fixture
async def plugin_app(engine: AsyncEngine, redis_client): async def plugin_app(engine: AsyncEngine, redis_client):
"""FastAPI app with plugin registry initialized for testing.""" """FastAPI app with plugin registry initialized for testing."""
@@ -82,7 +83,7 @@ async def plugin_client(plugin_app) -> AsyncClient:
@pytest_asyncio.fixture @pytest_asyncio.fixture
async def authed_plugin_client(plugin_client: AsyncClient, db_session: AsyncSession) -> AsyncClient: async def authed_plugin_client(plugin_client: AsyncClient, db_session: AsyncSession) -> AsyncClient:
"""Authenticated admin client for plugin tests.""" """Authenticated admin client for plugin tests."""
seed = await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(plugin_client, "admin@tenanta.com") await login_client(plugin_client, "admin@tenanta.com")
return plugin_client return plugin_client
@@ -98,6 +99,7 @@ async def db_session_for_plugins(engine: AsyncEngine) -> AsyncSession:
# ─── AC1: GET /api/v1/plugins → 200 + list of plugins with status ─── # ─── AC1: GET /api/v1/plugins → 200 + list of plugins with status ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac01_list_plugins(authed_plugin_client: AsyncClient): async def test_ac01_list_plugins(authed_plugin_client: AsyncClient):
"""AC1: GET /api/v1/plugins returns 200 with plugin list and status.""" """AC1: GET /api/v1/plugins returns 200 with plugin list and status."""
@@ -119,10 +121,15 @@ async def test_ac01_list_plugins(authed_plugin_client: AsyncClient):
# ─── AC2: POST /api/v1/plugins/{name}/install → 200, status=installed, migrations run ─── # ─── AC2: POST /api/v1/plugins/{name}/install → 200, status=installed, migrations run ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac02_install_plugin(authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession): async def test_ac02_install_plugin(
authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession
):
"""AC2: Install plugin runs migrations and sets status=installed.""" """AC2: Install plugin runs migrations and sets status=installed."""
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
data = resp.json() data = resp.json()
assert data["name"] == "test_sample" assert data["name"] == "test_sample"
@@ -132,22 +139,29 @@ async def test_ac02_install_plugin(authed_plugin_client: AsyncClient, db_session
# Verify migration table was created (tenant_id column exists) # Verify migration table was created (tenant_id column exists)
result = await db_session_for_plugins.execute( result = await db_session_for_plugins.execute(
text("SELECT column_name FROM information_schema.columns WHERE table_name = 'plugin_test_data' AND column_name = 'tenant_id'") text(
"SELECT column_name FROM information_schema.columns WHERE table_name = 'plugin_test_data' AND column_name = 'tenant_id'"
)
) )
assert result.fetchone() is not None, "plugin_test_data table should have tenant_id column" assert result.fetchone() is not None, "plugin_test_data table should have tenant_id column"
# ─── AC3: POST /api/v1/plugins/{name}/activate → 200, status=active, routes registered ─── # ─── AC3: POST /api/v1/plugins/{name}/activate → 200, status=active, routes registered ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac03_activate_plugin(authed_plugin_client: AsyncClient): async def test_ac03_activate_plugin(authed_plugin_client: AsyncClient):
"""AC3: Activate plugin sets status=active and registers event listeners.""" """AC3: Activate plugin sets status=active and registers event listeners."""
# Install first # Install first
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
# Activate # Activate
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
data = resp.json() data = resp.json()
assert data["name"] == "test_sample" assert data["name"] == "test_sample"
@@ -157,6 +171,7 @@ async def test_ac03_activate_plugin(authed_plugin_client: AsyncClient):
# ─── AC4: POST /api/v1/plugins/{name}/deactivate → 200, status=inactive, routes unregistered ─── # ─── AC4: POST /api/v1/plugins/{name}/deactivate → 200, status=inactive, routes unregistered ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac04_deactivate_plugin(authed_plugin_client: AsyncClient): async def test_ac04_deactivate_plugin(authed_plugin_client: AsyncClient):
"""AC4: Deactivate plugin sets status=inactive and unregisters event listeners.""" """AC4: Deactivate plugin sets status=inactive and unregisters event listeners."""
@@ -165,7 +180,9 @@ async def test_ac04_deactivate_plugin(authed_plugin_client: AsyncClient):
await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER) await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER)
# Deactivate # Deactivate
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
data = resp.json() data = resp.json()
assert data["name"] == "test_sample" assert data["name"] == "test_sample"
@@ -175,8 +192,11 @@ async def test_ac04_deactivate_plugin(authed_plugin_client: AsyncClient):
# ─── AC5: DELETE /api/v1/plugins/{name} → 200, plugin removed ─── # ─── AC5: DELETE /api/v1/plugins/{name} → 200, plugin removed ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac05_uninstall_plugin(authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession): async def test_ac05_uninstall_plugin(
authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession
):
"""AC5: Uninstall plugin removes DB record.""" """AC5: Uninstall plugin removes DB record."""
# Install first # Install first
await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER)
@@ -197,6 +217,7 @@ async def test_ac05_uninstall_plugin(authed_plugin_client: AsyncClient, db_sessi
# ─── AC6: DELETE /api/v1/plugins/{name}?remove_data=true → 200, plugin tables dropped ─── # ─── AC6: DELETE /api/v1/plugins/{name}?remove_data=true → 200, plugin tables dropped ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac06_uninstall_remove_data(authed_plugin_client: AsyncClient, engine: AsyncEngine): async def test_ac06_uninstall_remove_data(authed_plugin_client: AsyncClient, engine: AsyncEngine):
"""AC6: Uninstall with remove_data=true drops plugin tables.""" """AC6: Uninstall with remove_data=true drops plugin tables."""
@@ -207,6 +228,7 @@ async def test_ac06_uninstall_remove_data(authed_plugin_client: AsyncClient, eng
def _check_table(sync_conn): def _check_table(sync_conn):
insp = inspect(sync_conn) insp = inspect(sync_conn)
return "plugin_test_data" in insp.get_table_names() return "plugin_test_data" in insp.get_table_names()
async with engine.connect() as conn: async with engine.connect() as conn:
table_exists_before = await conn.run_sync(_check_table) table_exists_before = await conn.run_sync(_check_table)
assert table_exists_before, "plugin_test_data table should exist after install" assert table_exists_before, "plugin_test_data table should exist after install"
@@ -223,11 +245,14 @@ async def test_ac06_uninstall_remove_data(authed_plugin_client: AsyncClient, eng
# Verify table is gone # Verify table is gone
async with engine.connect() as conn: async with engine.connect() as conn:
table_exists_after = await conn.run_sync(_check_table) table_exists_after = await conn.run_sync(_check_table)
assert not table_exists_after, "plugin_test_data table should be dropped after uninstall with remove_data=true" assert (
not table_exists_after
), "plugin_test_data table should be dropped after uninstall with remove_data=true"
# ─── AC7: GET /api/v1/plugins/manifest → 200 + manifest schema documentation ─── # ─── AC7: GET /api/v1/plugins/manifest → 200 + manifest schema documentation ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac07_manifest_schema(authed_plugin_client: AsyncClient): async def test_ac07_manifest_schema(authed_plugin_client: AsyncClient):
"""AC7: GET /api/v1/plugins/manifest returns schema documentation.""" """AC7: GET /api/v1/plugins/manifest returns schema documentation."""
@@ -250,6 +275,7 @@ async def test_ac07_manifest_schema(authed_plugin_client: AsyncClient):
# ─── AC8: Plugin activation registers event listeners on event bus ─── # ─── AC8: Plugin activation registers event listeners on event bus ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac08_activation_registers_event_listeners( async def test_ac08_activation_registers_event_listeners(
authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession
@@ -281,6 +307,7 @@ async def test_ac08_activation_registers_event_listeners(
# ─── AC9: Plugin deactivation unregisters event listeners ─── # ─── AC9: Plugin deactivation unregisters event listeners ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac09_deactivation_unregisters_event_listeners(authed_plugin_client: AsyncClient): async def test_ac09_deactivation_unregisters_event_listeners(authed_plugin_client: AsyncClient):
"""AC9: Deactivating a plugin unregisters its event listeners from the event bus.""" """AC9: Deactivating a plugin unregisters its event listeners from the event bus."""
@@ -297,11 +324,15 @@ async def test_ac09_deactivation_unregisters_event_listeners(authed_plugin_clien
# Publish an event — handler should NOT be called after deactivation # Publish an event — handler should NOT be called after deactivation
initial_log_count = len(plugin.event_log) initial_log_count = len(plugin.event_log)
await event_bus.publish("company.created", {"company_id": "test-456", "name": "After Deactivate"}) await event_bus.publish(
"company.created", {"company_id": "test-456", "name": "After Deactivate"}
)
await asyncio.sleep(0.1) await asyncio.sleep(0.1)
# Event log should not have grown # Event log should not have grown
assert len(plugin.event_log) == initial_log_count, "Event handler should not be called after deactivation" assert (
len(plugin.event_log) == initial_log_count
), "Event handler should not be called after deactivation"
# Check event bus no longer has the plugin's handlers # Check event bus no longer has the plugin's handlers
handlers = event_bus._handlers.get("company.created", []) handlers = event_bus._handlers.get("company.created", [])
@@ -312,11 +343,16 @@ async def test_ac09_deactivation_unregisters_event_listeners(authed_plugin_clien
# ─── AC10: Plugin migration creates tables with tenant_id column ─── # ─── AC10: Plugin migration creates tables with tenant_id column ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac10_migration_creates_tenant_id(authed_plugin_client: AsyncClient, engine: AsyncEngine): async def test_ac10_migration_creates_tenant_id(
authed_plugin_client: AsyncClient, engine: AsyncEngine
):
"""AC10: Plugin migration creates tables that have tenant_id column.""" """AC10: Plugin migration creates tables that have tenant_id column."""
# Install the test_sample plugin which has a migration creating plugin_test_data # Install the test_sample plugin which has a migration creating plugin_test_data
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
# Verify the table has tenant_id column # Verify the table has tenant_id column
@@ -335,6 +371,7 @@ async def test_ac10_migration_creates_tenant_id(authed_plugin_client: AsyncClien
# ─── AC11: Plugin migration validator rejects tables without tenant_id ─── # ─── AC11: Plugin migration validator rejects tables without tenant_id ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac11_validator_rejects_no_tenant_id(authed_plugin_client: AsyncClient): async def test_ac11_validator_rejects_no_tenant_id(authed_plugin_client: AsyncClient):
"""AC11: Migration validator rejects tables created without tenant_id column.""" """AC11: Migration validator rejects tables created without tenant_id column."""
@@ -357,11 +394,16 @@ async def test_ac11_validator_rejects_no_tenant_id(authed_plugin_client: AsyncCl
# ─── AC12: Plugin DB migrations tracked in plugin_migrations table ─── # ─── AC12: Plugin DB migrations tracked in plugin_migrations table ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac12_migrations_tracked(authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession): async def test_ac12_migrations_tracked(
authed_plugin_client: AsyncClient, db_session_for_plugins: AsyncSession
):
"""AC12: Plugin migrations are tracked in plugin_migrations table.""" """AC12: Plugin migrations are tracked in plugin_migrations table."""
# Install the test_sample plugin # Install the test_sample plugin
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
# Check plugin_migrations table has a record # Check plugin_migrations table has a record
@@ -376,17 +418,22 @@ async def test_ac12_migrations_tracked(authed_plugin_client: AsyncClient, db_ses
# ─── AC13: Activating already-active plugin → idempotent (200, no error) ─── # ─── AC13: Activating already-active plugin → idempotent (200, no error) ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac13_activate_already_active(authed_plugin_client: AsyncClient): async def test_ac13_activate_already_active(authed_plugin_client: AsyncClient):
"""AC13: Activating an already-active plugin is idempotent (returns 200, no error).""" """AC13: Activating an already-active plugin is idempotent (returns 200, no error)."""
# Install and activate # Install and activate
await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER) await authed_plugin_client.post("/api/v1/plugins/test_sample/install", headers=ORIGIN_HEADER)
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
assert resp.json()["status"] == "active" assert resp.json()["status"] == "active"
# Activate again — should be idempotent # Activate again — should be idempotent
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
data = resp.json() data = resp.json()
assert data["status"] == "active" assert data["status"] == "active"
@@ -396,6 +443,7 @@ async def test_ac13_activate_already_active(authed_plugin_client: AsyncClient):
# ─── AC14: Deactivating inactive plugin → idempotent (200) ─── # ─── AC14: Deactivating inactive plugin → idempotent (200) ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_ac14_deactivate_already_inactive(authed_plugin_client: AsyncClient): async def test_ac14_deactivate_already_inactive(authed_plugin_client: AsyncClient):
"""AC14: Deactivating an already-inactive plugin is idempotent (returns 200, no error).""" """AC14: Deactivating an already-inactive plugin is idempotent (returns 200, no error)."""
@@ -404,12 +452,16 @@ async def test_ac14_deactivate_already_inactive(authed_plugin_client: AsyncClien
await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER) await authed_plugin_client.post("/api/v1/plugins/test_sample/activate", headers=ORIGIN_HEADER)
# Deactivate # Deactivate
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
assert resp.json()["status"] == "inactive" assert resp.json()["status"] == "inactive"
# Deactivate again — should be idempotent # Deactivate again — should be idempotent
resp = await authed_plugin_client.post("/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER) resp = await authed_plugin_client.post(
"/api/v1/plugins/test_sample/deactivate", headers=ORIGIN_HEADER
)
assert resp.status_code == 200 assert resp.status_code == 200
data = resp.json() data = resp.json()
assert data["status"] == "inactive" assert data["status"] == "inactive"
@@ -419,6 +471,7 @@ async def test_ac14_deactivate_already_inactive(authed_plugin_client: AsyncClien
# ─── Additional Tests: Direct MigrationRunner Unit Tests ─── # ─── Additional Tests: Direct MigrationRunner Unit Tests ───
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_discover_builtins(engine: AsyncEngine): async def test_registry_discover_builtins(engine: AsyncEngine):
"""Test that registry can discover built-in plugins.""" """Test that registry can discover built-in plugins."""
@@ -431,7 +484,9 @@ async def test_registry_discover_builtins(engine: AsyncEngine):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_list_plugins_mixed_states(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_list_plugins_mixed_states(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test listing plugins in various states (discovered, installed, active, inactive).""" """Test listing plugins in various states (discovered, installed, active, inactive)."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -464,7 +519,9 @@ async def test_registry_list_plugins_mixed_states(engine: AsyncEngine, db_sessio
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_install_idempotent(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_install_idempotent(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test that installing an already-installed plugin is idempotent.""" """Test that installing an already-installed plugin is idempotent."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -500,7 +557,9 @@ async def test_registry_not_found_errors(engine: AsyncEngine, db_session_for_plu
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_activate_without_install(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_activate_without_install(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test that activating without install raises error.""" """Test that activating without install raises error."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -511,7 +570,9 @@ async def test_registry_activate_without_install(engine: AsyncEngine, db_session
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_migration_runner_drop_tables(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_migration_runner_drop_tables(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test MigrationRunner.drop_plugin_tables removes tables and migration records.""" """Test MigrationRunner.drop_plugin_tables removes tables and migration records."""
runner = MigrationRunner(engine) runner = MigrationRunner(engine)
@@ -545,7 +606,9 @@ async def test_migration_runner_drop_tables(engine: AsyncEngine, db_session_for_
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_migration_runner_file_not_found(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_migration_runner_file_not_found(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test MigrationRunner raises FileNotFoundError for missing migration file.""" """Test MigrationRunner raises FileNotFoundError for missing migration file."""
runner = MigrationRunner(engine) runner = MigrationRunner(engine)
with pytest.raises(FileNotFoundError): with pytest.raises(FileNotFoundError):
@@ -592,7 +655,7 @@ def test_plugin_manifest_validation():
assert m2.name == "myplugin" assert m2.name == "myplugin"
# Invalid name with special chars # Invalid name with special chars
with pytest.raises(Exception): with pytest.raises(Exception): # noqa: B017
PluginManifest(name="my-plugin!", version="1.0.0", display_name="Test") PluginManifest(name="my-plugin!", version="1.0.0", display_name="Test")
@@ -606,6 +669,7 @@ def test_base_plugin_repr():
def test_base_plugin_no_manifest_error(): def test_base_plugin_no_manifest_error():
"""Test that BasePlugin without manifest raises ValueError.""" """Test that BasePlugin without manifest raises ValueError."""
class NoManifestPlugin(BasePlugin): class NoManifestPlugin(BasePlugin):
pass pass
@@ -623,10 +687,14 @@ def test_base_plugin_name_version_properties():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_base_plugin_on_install_default(): async def test_base_plugin_on_install_default():
"""Test default on_install returns None (no-op).""" """Test default on_install returns None (no-op)."""
class MinimalPlugin(BasePlugin): class MinimalPlugin(BasePlugin):
manifest = PluginManifest( manifest = PluginManifest(
name="minimal", version="1.0.0", display_name="Minimal", name="minimal",
version="1.0.0",
display_name="Minimal",
) )
plugin = MinimalPlugin() plugin = MinimalPlugin()
result = await plugin.on_install(None, None) result = await plugin.on_install(None, None)
assert result is None assert result is None
@@ -635,10 +703,14 @@ async def test_base_plugin_on_install_default():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_base_plugin_on_uninstall_default(): async def test_base_plugin_on_uninstall_default():
"""Test default on_uninstall returns None (no-op).""" """Test default on_uninstall returns None (no-op)."""
class MinimalPlugin(BasePlugin): class MinimalPlugin(BasePlugin):
manifest = PluginManifest( manifest = PluginManifest(
name="minimal2", version="1.0.0", display_name="Minimal", name="minimal2",
version="1.0.0",
display_name="Minimal",
) )
plugin = MinimalPlugin() plugin = MinimalPlugin()
result = await plugin.on_uninstall(None, None) result = await plugin.on_uninstall(None, None)
assert result is None assert result is None
@@ -653,11 +725,15 @@ def test_base_plugin_get_routes_empty():
def test_base_plugin_make_event_handler_noop(): def test_base_plugin_make_event_handler_noop():
"""Test that _make_event_handler creates noop for unknown events.""" """Test that _make_event_handler creates noop for unknown events."""
class MinimalPlugin(BasePlugin): class MinimalPlugin(BasePlugin):
manifest = PluginManifest( manifest = PluginManifest(
name="minimal3", version="1.0.0", display_name="Minimal", name="minimal3",
version="1.0.0",
display_name="Minimal",
events=["unknown.event"], events=["unknown.event"],
) )
plugin = MinimalPlugin() plugin = MinimalPlugin()
handler = plugin._make_event_handler("unknown.event") handler = plugin._make_event_handler("unknown.event")
assert handler is not None assert handler is not None
@@ -669,6 +745,7 @@ def test_base_plugin_make_event_handler_noop():
async def test_service_manifest_schema(): async def test_service_manifest_schema():
"""Test plugin service get_manifest_schema returns dict.""" """Test plugin service get_manifest_schema returns dict."""
from app.services.plugin_service import PluginService from app.services.plugin_service import PluginService
service = PluginService() service = PluginService()
schema = service.get_manifest_schema() schema = service.get_manifest_schema()
assert isinstance(schema, dict) assert isinstance(schema, dict)
@@ -694,10 +771,13 @@ async def test_uninstall_inactive_plugin(engine: AsyncEngine, db_session_for_plu
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_activate_with_app_routes(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_activate_with_app_routes(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test that activating a plugin with a FastAPI app registers routes.""" """Test that activating a plugin with a FastAPI app registers routes."""
from fastapi import FastAPI, APIRouter from fastapi import APIRouter, FastAPI
from app.plugins.manifest import PluginManifest, PluginRouteDef
from app.plugins.manifest import PluginManifest
# Create a plugin with a route # Create a plugin with a route
class RoutePlugin(BasePlugin): class RoutePlugin(BasePlugin):
@@ -707,11 +787,14 @@ async def test_registry_activate_with_app_routes(engine: AsyncEngine, db_session
display_name="Route Plugin", display_name="Route Plugin",
events=["test.event"], events=["test.event"],
) )
def get_routes(self) -> list: def get_routes(self) -> list:
test_router = APIRouter() test_router = APIRouter()
@test_router.get("/api/v1/route-plugin/test") @test_router.get("/api/v1/route-plugin/test")
async def test_endpoint(): async def test_endpoint():
return {"status": "ok"} return {"status": "ok"}
return [test_router] return [test_router]
app = FastAPI() app = FastAPI()
@@ -737,7 +820,9 @@ async def test_registry_activate_with_app_routes(engine: AsyncEngine, db_session
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_uninstall_with_remove_data(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_uninstall_with_remove_data(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test registry uninstall with remove_data drops tables.""" """Test registry uninstall with remove_data drops tables."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -763,7 +848,9 @@ async def test_registry_not_initialized_errors():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_deactivate_already_inactive_direct(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_deactivate_already_inactive_direct(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test registry.deactivate is idempotent when already inactive.""" """Test registry.deactivate is idempotent when already inactive."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -785,9 +872,12 @@ async def test_registry_deactivate_already_inactive_direct(engine: AsyncEngine,
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_plugin_service_install_error_handling(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_plugin_service_install_error_handling(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test plugin service raises ValueError for unknown plugin.""" """Test plugin service raises ValueError for unknown plugin."""
from app.services.plugin_service import PluginService from app.services.plugin_service import PluginService
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
service = PluginService(registry=registry) service = PluginService(registry=registry)
@@ -797,9 +887,12 @@ async def test_plugin_service_install_error_handling(engine: AsyncEngine, db_ses
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_plugin_service_activate_error_handling(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_plugin_service_activate_error_handling(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test plugin service raises ValueError when activating uninstalled plugin.""" """Test plugin service raises ValueError when activating uninstalled plugin."""
from app.services.plugin_service import PluginService from app.services.plugin_service import PluginService
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
registry.register_plugin(TestSamplePlugin()) registry.register_plugin(TestSamplePlugin())
@@ -810,9 +903,12 @@ async def test_plugin_service_activate_error_handling(engine: AsyncEngine, db_se
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_plugin_service_deactivate_error_handling(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_plugin_service_deactivate_error_handling(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test plugin service raises ValueError for deactivating unknown plugin.""" """Test plugin service raises ValueError for deactivating unknown plugin."""
from app.services.plugin_service import PluginService from app.services.plugin_service import PluginService
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
service = PluginService(registry=registry) service = PluginService(registry=registry)
@@ -822,9 +918,12 @@ async def test_plugin_service_deactivate_error_handling(engine: AsyncEngine, db_
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_plugin_service_uninstall_error_handling(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_plugin_service_uninstall_error_handling(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test plugin service raises ValueError for uninstalling unknown plugin.""" """Test plugin service raises ValueError for uninstalling unknown plugin."""
from app.services.plugin_service import PluginService from app.services.plugin_service import PluginService
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
service = PluginService(registry=registry) service = PluginService(registry=registry)
@@ -842,7 +941,9 @@ async def test_migration_runner_split_sql_dollar_quotes():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_registry_list_plugins_db_only_record(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_registry_list_plugins_db_only_record(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test listing plugins includes DB-only records (installed but not discovered).""" """Test listing plugins includes DB-only records (installed but not discovered)."""
registry = reset_registry_for_testing() registry = reset_registry_for_testing()
registry.initialize(engine, None) registry.initialize(engine, None)
@@ -869,11 +970,15 @@ async def test_registry_list_plugins_db_only_record(engine: AsyncEngine, db_sess
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_base_plugin_on_event_fallback(): async def test_base_plugin_on_event_fallback():
"""Test that on_event fallback handler works.""" """Test that on_event fallback handler works."""
class FallbackPlugin(BasePlugin): class FallbackPlugin(BasePlugin):
manifest = PluginManifest( manifest = PluginManifest(
name="fallback", version="1.0.0", display_name="Fallback", name="fallback",
version="1.0.0",
display_name="Fallback",
events=["custom.event"], events=["custom.event"],
) )
async def on_event(self, payload: dict[str, Any]) -> None: async def on_event(self, payload: dict[str, Any]) -> None:
self.event_received = payload self.event_received = payload
@@ -884,7 +989,9 @@ async def test_base_plugin_on_event_fallback():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_migration_runner_valid_migration(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_migration_runner_valid_migration(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test MigrationRunner directly with a valid migration (has tenant_id).""" """Test MigrationRunner directly with a valid migration (has tenant_id)."""
runner = MigrationRunner(engine) runner = MigrationRunner(engine)
record = await runner.run_migration( record = await runner.run_migration(
@@ -899,7 +1006,9 @@ async def test_migration_runner_valid_migration(engine: AsyncEngine, db_session_
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_migration_runner_invalid_migration(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_migration_runner_invalid_migration(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test MigrationRunner directly rejects migration without tenant_id.""" """Test MigrationRunner directly rejects migration without tenant_id."""
runner = MigrationRunner(engine) runner = MigrationRunner(engine)
with pytest.raises(MigrationValidationError) as exc_info: with pytest.raises(MigrationValidationError) as exc_info:
@@ -912,7 +1021,9 @@ async def test_migration_runner_invalid_migration(engine: AsyncEngine, db_sessio
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_migration_runner_idempotent(engine: AsyncEngine, db_session_for_plugins: AsyncSession): async def test_migration_runner_idempotent(
engine: AsyncEngine, db_session_for_plugins: AsyncSession
):
"""Test MigrationRunner skips already-applied migrations.""" """Test MigrationRunner skips already-applied migrations."""
runner = MigrationRunner(engine) runner = MigrationRunner(engine)
# Run migration # Run migration
+40 -17
View File
@@ -7,16 +7,15 @@ import uuid
import pytest import pytest
import pytest_asyncio import pytest_asyncio
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession
from sqlalchemy.ext.asyncio import AsyncEngine, async_sessionmaker, AsyncSession
from app.core.db import Base, reset_engine_for_testing, close_engine from app.core.db import close_engine, reset_engine_for_testing
from app.core.service_container import get_container from app.core.service_container import get_container
from app.main import create_app from app.main import create_app
from app.plugins.registry import get_registry, reset_registry_for_testing
from app.plugins.builtins.tags import TagsPlugin from app.plugins.builtins.tags import TagsPlugin
from app.plugins.registry import reset_registry_for_testing
from app.services.plugin_service import reset_plugin_service_for_testing from app.services.plugin_service import reset_plugin_service_for_testing
from tests.conftest import seed_tenant_and_users, login_client, ORIGIN_HEADER from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest_asyncio.fixture @pytest_asyncio.fixture
@@ -48,7 +47,7 @@ async def plugin_client(plugin_app) -> AsyncClient:
@pytest_asyncio.fixture @pytest_asyncio.fixture
async def authed_client(plugin_client: AsyncClient, db_session: AsyncSession) -> AsyncClient: async def authed_client(plugin_client: AsyncClient, db_session: AsyncSession) -> AsyncClient:
"""Authenticated admin client with seeded data.""" """Authenticated admin client with seeded data."""
seed = await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(plugin_client, "admin@tenanta.com") await login_client(plugin_client, "admin@tenanta.com")
# Install + activate the tags plugin # Install + activate the tags plugin
resp = await plugin_client.post("/api/v1/plugins/tags/install", headers=ORIGIN_HEADER) resp = await plugin_client.post("/api/v1/plugins/tags/install", headers=ORIGIN_HEADER)
@@ -91,7 +90,7 @@ async def test_list_tags_with_counts(authed_client: AsyncClient):
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
assert resp.status_code == 201 assert resp.status_code == 201
tag_id = resp.json()["id"] resp.json()["id"]
# List tags — should show entity_count=0 # List tags — should show entity_count=0
resp = await authed_client.get("/api/v1/tags", headers=ORIGIN_HEADER) resp = await authed_client.get("/api/v1/tags", headers=ORIGIN_HEADER)
@@ -302,11 +301,15 @@ async def test_update_tag_invalid_id(authed_client: AsyncClient):
async def test_update_tag_duplicate_name(authed_client: AsyncClient): async def test_update_tag_duplicate_name(authed_client: AsyncClient):
"""PATCH /api/v1/tags/{id} with existing name → 409.""" """PATCH /api/v1/tags/{id} with existing name → 409."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "TagA", "color": "#AAAAAA"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "TagA", "color": "#AAAAAA"},
headers=ORIGIN_HEADER,
) )
assert resp.status_code == 201 assert resp.status_code == 201
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "TagB", "color": "#BBBBBB"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "TagB", "color": "#BBBBBB"},
headers=ORIGIN_HEADER,
) )
assert resp.status_code == 201 assert resp.status_code == 201
tag_b_id = resp.json()["id"] tag_b_id = resp.json()["id"]
@@ -322,7 +325,9 @@ async def test_update_tag_duplicate_name(authed_client: AsyncClient):
async def test_update_tag_color_only(authed_client: AsyncClient): async def test_update_tag_color_only(authed_client: AsyncClient):
"""PATCH /api/v1/tags/{id} with only color → 200.""" """PATCH /api/v1/tags/{id} with only color → 200."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "ColorTag", "color": "#CCCCCC"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "ColorTag", "color": "#CCCCCC"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
resp = await authed_client.patch( resp = await authed_client.patch(
@@ -352,7 +357,9 @@ async def test_delete_tag_invalid_id(authed_client: AsyncClient):
async def test_assign_tag_invalid_entity_type(authed_client: AsyncClient): async def test_assign_tag_invalid_entity_type(authed_client: AsyncClient):
"""POST /api/v1/tags/assign with invalid entity_type → 400.""" """POST /api/v1/tags/assign with invalid entity_type → 400."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "ETTag", "color": "#EEEEEE"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "ETTag", "color": "#EEEEEE"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
resp = await authed_client.post( resp = await authed_client.post(
@@ -368,7 +375,11 @@ async def test_assign_tag_not_found(authed_client: AsyncClient):
"""POST /api/v1/tags/assign with nonexistent tag → 404.""" """POST /api/v1/tags/assign with nonexistent tag → 404."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags/assign", "/api/v1/tags/assign",
json={"tag_id": str(uuid.uuid4()), "entity_type": "company", "entity_id": str(uuid.uuid4())}, json={
"tag_id": str(uuid.uuid4()),
"entity_type": "company",
"entity_id": str(uuid.uuid4()),
},
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
assert resp.status_code == 404 assert resp.status_code == 404
@@ -378,7 +389,9 @@ async def test_assign_tag_not_found(authed_client: AsyncClient):
async def test_assign_tag_already_assigned(authed_client: AsyncClient): async def test_assign_tag_already_assigned(authed_client: AsyncClient):
"""POST /api/v1/tags/assign twice → already_assigned=True.""" """POST /api/v1/tags/assign twice → already_assigned=True."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "DupAssign", "color": "#FFFFFF"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "DupAssign", "color": "#FFFFFF"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
entity_id = str(uuid.uuid4()) entity_id = str(uuid.uuid4())
@@ -413,7 +426,9 @@ async def test_assign_tag_invalid_ids(authed_client: AsyncClient):
async def test_unassign_tag_not_found(authed_client: AsyncClient): async def test_unassign_tag_not_found(authed_client: AsyncClient):
"""DELETE /api/v1/tags/assign with nonexistent assignment → 404.""" """DELETE /api/v1/tags/assign with nonexistent assignment → 404."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "Unassign404", "color": "#ABABAB"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "Unassign404", "color": "#ABABAB"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
resp = await authed_client.request( resp = await authed_client.request(
@@ -429,7 +444,9 @@ async def test_unassign_tag_not_found(authed_client: AsyncClient):
async def test_bulk_assign_invalid_entity_type(authed_client: AsyncClient): async def test_bulk_assign_invalid_entity_type(authed_client: AsyncClient):
"""POST /api/v1/tags/bulk-assign with invalid entity_type → 400.""" """POST /api/v1/tags/bulk-assign with invalid entity_type → 400."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "BulkBad", "color": "#000001"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "BulkBad", "color": "#000001"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
resp = await authed_client.post( resp = await authed_client.post(
@@ -445,7 +462,11 @@ async def test_bulk_assign_tag_not_found(authed_client: AsyncClient):
"""POST /api/v1/tags/bulk-assign with nonexistent tag → 404.""" """POST /api/v1/tags/bulk-assign with nonexistent tag → 404."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags/bulk-assign", "/api/v1/tags/bulk-assign",
json={"tag_ids": [str(uuid.uuid4())], "entity_type": "file", "entity_id": str(uuid.uuid4())}, json={
"tag_ids": [str(uuid.uuid4())],
"entity_type": "file",
"entity_id": str(uuid.uuid4()),
},
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
) )
assert resp.status_code == 404 assert resp.status_code == 404
@@ -455,7 +476,9 @@ async def test_bulk_assign_tag_not_found(authed_client: AsyncClient):
async def test_bulk_assign_already_assigned(authed_client: AsyncClient): async def test_bulk_assign_already_assigned(authed_client: AsyncClient):
"""POST /api/v1/tags/bulk-assign twice → already_assigned populated.""" """POST /api/v1/tags/bulk-assign twice → already_assigned populated."""
resp = await authed_client.post( resp = await authed_client.post(
"/api/v1/tags", json={"name": "BulkDup1", "color": "#001100"}, headers=ORIGIN_HEADER, "/api/v1/tags",
json={"name": "BulkDup1", "color": "#001100"},
headers=ORIGIN_HEADER,
) )
tag_id = resp.json()["id"] tag_id = resp.json()["id"]
entity_id = str(uuid.uuid4()) entity_id = str(uuid.uuid4())
+14 -5
View File
@@ -5,12 +5,10 @@ from __future__ import annotations
import pytest import pytest
from httpx import AsyncClient from httpx import AsyncClient
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client
from app.models.audit import AuditLog from app.models.audit import AuditLog
from app.models.notification import Notification from app.models.notification import Notification
from app.models.company import Company from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -60,6 +58,7 @@ class TestUserManagement:
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
# Get viewer user ID # Get viewer user ID
from app.models.user import User from app.models.user import User
q = select(User).where(User.email == "viewer@tenanta.com") q = select(User).where(User.email == "viewer@tenanta.com")
result = await db_session.execute(q) result = await db_session.execute(q)
viewer = result.scalar_one() viewer = result.scalar_one()
@@ -76,6 +75,7 @@ class TestUserManagement:
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
from app.models.user import User from app.models.user import User
q = select(User).where(User.email == "viewer@tenanta.com") q = select(User).where(User.email == "viewer@tenanta.com")
result = await db_session.execute(q) result = await db_session.execute(q)
viewer = result.scalar_one() viewer = result.scalar_one()
@@ -102,7 +102,9 @@ class TestRoleManagement:
assert "permissions" in data["items"][0] assert "permissions" in data["items"][0]
assert "field_permissions" in data["items"][0] assert "field_permissions" in data["items"][0]
async def test_create_role_with_custom_permissions_returns_201(self, client: AsyncClient, db_session): async def test_create_role_with_custom_permissions_returns_201(
self, client: AsyncClient, db_session
):
"""AC 16: POST /api/v1/roles with custom permissions -> 201.""" """AC 16: POST /api/v1/roles with custom permissions -> 201."""
await seed_tenant_and_users(db_session) await seed_tenant_and_users(db_session)
await login_client(client, "admin@tenanta.com") await login_client(client, "admin@tenanta.com")
@@ -110,7 +112,9 @@ class TestRoleManagement:
"/api/v1/roles", "/api/v1/roles",
json={ json={
"name": "manager", "name": "manager",
"permissions": {"companies": {"read": True, "create": True, "update": True, "delete": True}}, "permissions": {
"companies": {"read": True, "create": True, "update": True, "delete": True}
},
"field_permissions": {"annual_revenue": "read"}, "field_permissions": {"annual_revenue": "read"},
}, },
headers=ORIGIN_HEADER, headers=ORIGIN_HEADER,
@@ -172,6 +176,7 @@ class TestFieldPermissions:
# Create a user with sales_rep role # Create a user with sales_rep role
from app.core.auth import hash_password from app.core.auth import hash_password
from app.models.user import User, UserTenant from app.models.user import User, UserTenant
sales_user = User( sales_user = User(
tenant_id=seed["tenant_a"].id, tenant_id=seed["tenant_a"].id,
email="sales@tenanta.com", email="sales@tenanta.com",
@@ -213,6 +218,7 @@ class TestAuditLog:
# Check audit_log table # Check audit_log table
from sqlalchemy import select as sel from sqlalchemy import select as sel
q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "create") q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "create")
result = await db_session.execute(q) result = await db_session.execute(q)
entries = result.scalars().all() entries = result.scalars().all()
@@ -230,6 +236,7 @@ class TestAuditLog:
assert resp.status_code == 200 assert resp.status_code == 200
from sqlalchemy import select as sel from sqlalchemy import select as sel
q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "update") q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "update")
result = await db_session.execute(q) result = await db_session.execute(q)
entries = result.scalars().all() entries = result.scalars().all()
@@ -246,6 +253,7 @@ class TestAuditLog:
assert resp.status_code == 204 assert resp.status_code == 204
from sqlalchemy import select as sel from sqlalchemy import select as sel
q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "delete") q = sel(AuditLog).where(AuditLog.entity_type == "company", AuditLog.action == "delete")
result = await db_session.execute(q) result = await db_session.execute(q)
entries = result.scalars().all() entries = result.scalars().all()
@@ -275,6 +283,7 @@ class TestNotificationOnAssign:
# Check notification was created for the new user # Check notification was created for the new user
from sqlalchemy import select as sel from sqlalchemy import select as sel
q = sel(Notification).where(Notification.user_id == new_user_id) q = sel(Notification).where(Notification.user_id == new_user_id)
result = await db_session.execute(q) result = await db_session.execute(q)
notifs = result.scalars().all() notifs = result.scalars().all()
+456 -245
View File
File diff suppressed because it is too large Load Diff