chore: fix all ruff lint errors + format — 0 errors, 306 tests pass
This commit is contained in:
+12
-3
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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,10 +43,14 @@ 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 = (
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Notification)
|
||||||
|
.where(
|
||||||
Notification.tenant_id == tenant_id,
|
Notification.tenant_id == tenant_id,
|
||||||
Notification.user_id == user_id,
|
Notification.user_id == user_id,
|
||||||
)
|
)
|
||||||
|
)
|
||||||
total = (await db.execute(count_q)).scalar() or 0
|
total = (await db.execute(count_q)).scalar() or 0
|
||||||
|
|
||||||
# Query — unread first (read_at IS NULL), then newest
|
# Query — unread first (read_at IS NULL), then newest
|
||||||
@@ -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,11 +108,15 @@ 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 = (
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Notification)
|
||||||
|
.where(
|
||||||
Notification.tenant_id == tenant_id,
|
Notification.tenant_id == tenant_id,
|
||||||
Notification.user_id == user_id,
|
Notification.user_id == user_id,
|
||||||
Notification.read_at.is_(None),
|
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
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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"),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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
@@ -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}>"
|
||||||
|
)
|
||||||
|
|||||||
@@ -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"]
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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),
|
||||||
|
|||||||
@@ -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"),
|
||||||
|
|||||||
@@ -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,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
@@ -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",
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
+19
-15
@@ -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,7 +309,8 @@ 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,
|
"name": name,
|
||||||
"display_name": record.display_name,
|
"display_name": record.display_name,
|
||||||
"version": record.version,
|
"version": record.version,
|
||||||
@@ -321,9 +322,11 @@ class PluginRegistry:
|
|||||||
"events": plugin.manifest.events,
|
"events": plugin.manifest.events,
|
||||||
"migrations": plugin.manifest.migrations,
|
"migrations": plugin.manifest.migrations,
|
||||||
"permissions": plugin.manifest.permissions,
|
"permissions": plugin.manifest.permissions,
|
||||||
})
|
}
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
plugins_list.append({
|
plugins_list.append(
|
||||||
|
{
|
||||||
"name": name,
|
"name": name,
|
||||||
"display_name": plugin.manifest.display_name,
|
"display_name": plugin.manifest.display_name,
|
||||||
"version": plugin.version,
|
"version": plugin.version,
|
||||||
@@ -335,12 +338,14 @@ class PluginRegistry:
|
|||||||
"events": plugin.manifest.events,
|
"events": plugin.manifest.events,
|
||||||
"migrations": plugin.manifest.migrations,
|
"migrations": plugin.manifest.migrations,
|
||||||
"permissions": plugin.manifest.permissions,
|
"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,
|
"name": name,
|
||||||
"display_name": record.display_name,
|
"display_name": record.display_name,
|
||||||
"version": record.version,
|
"version": record.version,
|
||||||
@@ -352,7 +357,8 @@ class PluginRegistry:
|
|||||||
"events": [],
|
"events": [],
|
||||||
"migrations": [],
|
"migrations": [],
|
||||||
"permissions": [],
|
"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
@@ -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
|
||||||
|
)
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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,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
|
||||||
|
|||||||
@@ -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,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
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
@@ -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
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
@@ -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 = (
|
||||||
|
select(AIMessage)
|
||||||
|
.where(
|
||||||
AIMessage.conversation_id == conv.id,
|
AIMessage.conversation_id == conv.id,
|
||||||
AIMessage.tenant_id == tenant_id,
|
AIMessage.tenant_id == tenant_id,
|
||||||
).order_by(AIMessage.message_index)
|
)
|
||||||
|
.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),
|
**_conversation_to_dict(conv),
|
||||||
"messages": [_message_to_dict(m) for m in messages],
|
"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":
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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,7 +129,8 @@ 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),
|
"id": str(contact.id),
|
||||||
"first_name": contact.first_name,
|
"first_name": contact.first_name,
|
||||||
"last_name": contact.last_name,
|
"last_name": contact.last_name,
|
||||||
@@ -140,7 +139,8 @@ async def get_company_detail(
|
|||||||
"position": contact.position,
|
"position": contact.position,
|
||||||
"role_at_company": link.role_at_company,
|
"role_at_company": link.role_at_company,
|
||||||
"is_primary": link.is_primary,
|
"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)
|
||||||
|
|||||||
@@ -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),
|
"id": str(company.id),
|
||||||
"name": company.name,
|
"name": company.name,
|
||||||
"industry": company.industry,
|
"industry": company.industry,
|
||||||
"role_at_company": link.role_at_company,
|
"role_at_company": link.role_at_company,
|
||||||
"is_primary": link.is_primary,
|
"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
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
select(Contact)
|
||||||
|
.where(
|
||||||
Contact.tenant_id == tenant_id,
|
Contact.tenant_id == tenant_id,
|
||||||
Contact.deleted_at.is_(None),
|
Contact.deleted_at.is_(None),
|
||||||
).order_by(Contact.last_name, Contact.first_name)
|
)
|
||||||
|
.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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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 = (
|
||||||
|
select(WorkflowStepHistory)
|
||||||
|
.where(
|
||||||
WorkflowStepHistory.instance_id == inst_uuid,
|
WorkflowStepHistory.instance_id == inst_uuid,
|
||||||
WorkflowStepHistory.tenant_id == tenant_id,
|
WorkflowStepHistory.tenant_id == tenant_id,
|
||||||
).order_by(WorkflowStepHistory.created_at)
|
)
|
||||||
|
.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,
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
|
||||||
|
|||||||
+244
-44
@@ -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",
|
"message": "Here are actions",
|
||||||
"proposed_actions": [{"method": "GET", "path": "/api/v1/companies"}],
|
"proposed_actions": [{"method": "GET", "path": "/api/v1/companies"}],
|
||||||
"confidence": 0.95,
|
"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
@@ -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()
|
||||||
|
|||||||
@@ -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
@@ -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(
|
||||||
|
|||||||
@@ -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(
|
||||||
|
"company.deleted",
|
||||||
|
{
|
||||||
"entity_id": str(company_id),
|
"entity_id": str(company_id),
|
||||||
"company_id": str(company_id),
|
"company_id": str(company_id),
|
||||||
"tenant_id": str(tenant_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(
|
||||||
|
"contact.deleted",
|
||||||
|
{
|
||||||
"entity_id": str(contact_id),
|
"entity_id": str(contact_id),
|
||||||
"contact_id": str(contact_id),
|
"contact_id": str(contact_id),
|
||||||
"tenant_id": str(tenant_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
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||||
|
|||||||
+316
-105
@@ -3,24 +3,28 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime, timezone, timedelta
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from sqlalchemy import select
|
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
|
||||||
|
|
||||||
from httpx import AsyncClient
|
from httpx import AsyncClient
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
from tests.conftest import ORIGIN_HEADER, seed_tenant_and_users, login_client
|
|
||||||
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
|
|
||||||
from app.models.audit import AuditLog
|
|
||||||
from app.models.notification import Notification
|
from app.models.notification import Notification
|
||||||
|
from app.models.workflow import Workflow, WorkflowInstance, WorkflowStepHistory
|
||||||
from app.services.workflow_service import (
|
from app.services.workflow_service import (
|
||||||
create_workflow, list_workflows, get_workflow, update_workflow, delete_workflow,
|
advance_instance,
|
||||||
create_instance, list_instances, get_instance, advance_instance, cancel_instance,
|
auto_reject_timeout,
|
||||||
check_timeout, auto_reject_timeout, find_workflows_for_event, start_instance_for_event,
|
check_timeout,
|
||||||
|
create_instance,
|
||||||
|
create_workflow,
|
||||||
|
delete_workflow,
|
||||||
|
find_workflows_for_event,
|
||||||
|
get_instance,
|
||||||
|
list_workflows,
|
||||||
|
start_instance_for_event,
|
||||||
|
update_workflow,
|
||||||
)
|
)
|
||||||
|
from tests.conftest import ORIGIN_HEADER, login_client, seed_tenant_and_users
|
||||||
|
|
||||||
# ─── Workflow CRUD (ACs 8-12) ───
|
# ─── Workflow CRUD (ACs 8-12) ───
|
||||||
|
|
||||||
@@ -144,6 +148,7 @@ async def test_ac12_delete_workflow(client: AsyncClient, db_session):
|
|||||||
|
|
||||||
# ─── Instance Lifecycle (ACs 13-18) ───
|
# ─── Instance Lifecycle (ACs 13-18) ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ac13_create_instance(client: AsyncClient, db_session):
|
async def test_ac13_create_instance(client: AsyncClient, db_session):
|
||||||
"""AC13: POST /api/v1/workflows/{id}/instances returns 201, instance created with status=pending."""
|
"""AC13: POST /api/v1/workflows/{id}/instances returns 201, instance created with status=pending."""
|
||||||
@@ -328,6 +333,7 @@ async def test_ac18_cancel_instance(client: AsyncClient, db_session):
|
|||||||
|
|
||||||
# ─── Event-Triggered & Code-Engine (ACs 19-22) ───
|
# ─── Event-Triggered & Code-Engine (ACs 19-22) ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ac19_event_triggered_workflow(db_session):
|
async def test_ac19_event_triggered_workflow(db_session):
|
||||||
"""AC19: Event-triggered workflow — publish event → workflow instance auto-starts."""
|
"""AC19: Event-triggered workflow — publish event → workflow instance auto-starts."""
|
||||||
@@ -336,8 +342,10 @@ async def test_ac19_event_triggered_workflow(db_session):
|
|||||||
admin_id = seed["admin_a"].id
|
admin_id = seed["admin_a"].id
|
||||||
|
|
||||||
# Create a workflow with trigger_event
|
# Create a workflow with trigger_event
|
||||||
wf = await create_workflow(
|
await create_workflow(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
{
|
{
|
||||||
"name": "Event Triggered WF",
|
"name": "Event Triggered WF",
|
||||||
"trigger_event": "company.created",
|
"trigger_event": "company.created",
|
||||||
@@ -347,7 +355,9 @@ async def test_ac19_event_triggered_workflow(db_session):
|
|||||||
|
|
||||||
# Simulate event by calling start_instance_for_event
|
# Simulate event by calling start_instance_for_event
|
||||||
instances = await start_instance_for_event(
|
instances = await start_instance_for_event(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
"company.created",
|
"company.created",
|
||||||
context={"company_id": "test"},
|
context={"company_id": "test"},
|
||||||
)
|
)
|
||||||
@@ -363,7 +373,9 @@ async def test_ac20_step_history_created(db_session):
|
|||||||
admin_id = seed["admin_a"].id
|
admin_id = seed["admin_a"].id
|
||||||
|
|
||||||
wf = await create_workflow(
|
wf = await create_workflow(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
{"name": "History WF", "steps": VALID_STEPS},
|
{"name": "History WF", "steps": VALID_STEPS},
|
||||||
)
|
)
|
||||||
wf_id = wf["id"]
|
wf_id = wf["id"]
|
||||||
@@ -376,9 +388,7 @@ async def test_ac20_step_history_created(db_session):
|
|||||||
|
|
||||||
# Check step history
|
# Check step history
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
select(WorkflowStepHistory).where(
|
select(WorkflowStepHistory).where(WorkflowStepHistory.instance_id == uuid.UUID(inst_id))
|
||||||
WorkflowStepHistory.instance_id == uuid.UUID(inst_id)
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
history = result.scalars().all()
|
history = result.scalars().all()
|
||||||
assert len(history) >= 3 # entered + approved + entered (next step)
|
assert len(history) >= 3 # entered + approved + entered (next step)
|
||||||
@@ -387,7 +397,7 @@ async def test_ac20_step_history_created(db_session):
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_ac21_onboarding_workflow_on_user_creation(db_session):
|
async def test_ac21_onboarding_workflow_on_user_creation(db_session):
|
||||||
"""AC21: Code-engine workflow — onboarding workflow runs on user creation."""
|
"""AC21: Code-engine workflow — onboarding workflow runs on user creation."""
|
||||||
from app.workflows.code.onboarding import trigger_onboarding, get_onboarding_workflow_definition
|
from app.workflows.code.onboarding import trigger_onboarding
|
||||||
|
|
||||||
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
|
||||||
@@ -419,7 +429,9 @@ async def test_ac22_approval_timeout_auto_reject(db_session):
|
|||||||
|
|
||||||
# Create workflow with approval step
|
# Create workflow with approval step
|
||||||
wf = await create_workflow(
|
wf = await create_workflow(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
{
|
{
|
||||||
"name": "Timeout WF",
|
"name": "Timeout WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
@@ -430,7 +442,9 @@ async def test_ac22_approval_timeout_auto_reject(db_session):
|
|||||||
|
|
||||||
# Create instance with timeout_hours=1
|
# Create instance with timeout_hours=1
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
wf["id"],
|
wf["id"],
|
||||||
timeout_hours=1,
|
timeout_hours=1,
|
||||||
)
|
)
|
||||||
@@ -441,7 +455,7 @@ async def test_ac22_approval_timeout_auto_reject(db_session):
|
|||||||
select(WorkflowInstance).where(WorkflowInstance.id == inst_id)
|
select(WorkflowInstance).where(WorkflowInstance.id == inst_id)
|
||||||
)
|
)
|
||||||
instance = result.scalar_one()
|
instance = result.scalar_one()
|
||||||
instance.timeout_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
instance.timeout_at = datetime.now(UTC) - timedelta(hours=1)
|
||||||
await db_session.flush()
|
await db_session.flush()
|
||||||
|
|
||||||
# Check timeout detection
|
# Check timeout detection
|
||||||
@@ -461,6 +475,7 @@ async def test_ac22_approval_timeout_auto_reject(db_session):
|
|||||||
|
|
||||||
# ─── Edge Cases ───
|
# ─── Edge Cases ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_workflow_not_found(client: AsyncClient, db_session):
|
async def test_workflow_not_found(client: AsyncClient, db_session):
|
||||||
"""Edge case: Get non-existent workflow returns 404."""
|
"""Edge case: Get non-existent workflow returns 404."""
|
||||||
@@ -489,7 +504,10 @@ async def test_cancel_completed_instance_fails(client: AsyncClient, db_session):
|
|||||||
|
|
||||||
create_resp = await client.post(
|
create_resp = await client.post(
|
||||||
"/api/v1/workflows",
|
"/api/v1/workflows",
|
||||||
json={"name": "Complete WF", "steps": [{"name": "Only step", "type": "action", "config": {"action_type": "noop"}}]},
|
json={
|
||||||
|
"name": "Complete WF",
|
||||||
|
"steps": [{"name": "Only step", "type": "action", "config": {"action_type": "noop"}}],
|
||||||
|
},
|
||||||
headers=ORIGIN_HEADER,
|
headers=ORIGIN_HEADER,
|
||||||
)
|
)
|
||||||
wf_id = create_resp.json()["id"]
|
wf_id = create_resp.json()["id"]
|
||||||
@@ -530,8 +548,11 @@ async def test_workflow_tenant_isolation(client: AsyncClient, db_session):
|
|||||||
wf_id = create_resp.json()["id"]
|
wf_id = create_resp.json()["id"]
|
||||||
|
|
||||||
# Login as tenant B
|
# Login as tenant B
|
||||||
from httpx import ASGITransport, AsyncClient as AC
|
from httpx import ASGITransport
|
||||||
|
from httpx import AsyncClient as AC # noqa: N817
|
||||||
|
|
||||||
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:
|
||||||
@@ -557,6 +578,7 @@ async def test_workflow_create_rbac_viewer_blocked(client: AsyncClient, db_sessi
|
|||||||
|
|
||||||
# ─── WorkflowEngine Unit Tests (engine.py direct coverage) ───
|
# ─── WorkflowEngine Unit Tests (engine.py direct coverage) ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_engine_process_action_step_noop(db_session):
|
async def test_engine_process_action_step_noop(db_session):
|
||||||
"""Engine: action step with noop config advances to next step."""
|
"""Engine: action step with noop config advances to next step."""
|
||||||
@@ -566,13 +588,18 @@ async def test_engine_process_action_step_noop(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Engine Noop WF",
|
"name": "Engine Noop WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{"name": "Action", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Action", "type": "action", "config": {"action_type": "noop"}},
|
||||||
{"name": "Approval", "type": "approval", "config": {}},
|
{"name": "Approval", "type": "approval", "config": {}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
# Fetch the instance object
|
# Fetch the instance object
|
||||||
@@ -596,7 +623,11 @@ async def test_engine_process_action_step_create_notification(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Engine Notif WF",
|
"name": "Engine Notif WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -611,7 +642,8 @@ async def test_engine_process_action_step_create_notification(db_session):
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -638,10 +670,15 @@ async def test_engine_process_action_completes_last_step(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Single Step WF",
|
"name": "Single Step WF",
|
||||||
"steps": [{"name": "Only", "type": "action", "config": {"action_type": "noop"}}],
|
"steps": [{"name": "Only", "type": "action", "config": {"action_type": "noop"}}],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -664,10 +701,15 @@ async def test_engine_process_approval_step(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Approval WF",
|
"name": "Approval WF",
|
||||||
"steps": [{"name": "Approval", "type": "approval", "config": {}}],
|
"steps": [{"name": "Approval", "type": "approval", "config": {}}],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -690,7 +732,11 @@ async def test_engine_process_notification_step(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Notif Step WF",
|
"name": "Notif Step WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -704,7 +750,8 @@ async def test_engine_process_notification_step(db_session):
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -730,7 +777,11 @@ async def test_engine_process_condition_eq_true(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond Eq WF",
|
"name": "Cond Eq WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -747,9 +798,13 @@ async def test_engine_process_condition_eq_true(db_session):
|
|||||||
{"name": "False Step", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "False Step", "type": "action", "config": {"action_type": "noop"}},
|
||||||
{"name": "True Step", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "True Step", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"status": "active"},
|
context={"status": "active"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -772,7 +827,11 @@ async def test_engine_process_condition_eq_false(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond Eq False WF",
|
"name": "Cond Eq False WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -789,9 +848,13 @@ async def test_engine_process_condition_eq_false(db_session):
|
|||||||
{"name": "False Step", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "False Step", "type": "action", "config": {"action_type": "noop"}},
|
||||||
{"name": "True Step", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "True Step", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"status": "inactive"},
|
context={"status": "inactive"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -814,7 +877,11 @@ async def test_engine_process_condition_ne(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond NE WF",
|
"name": "Cond NE WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -828,9 +895,13 @@ async def test_engine_process_condition_ne(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"level": "high"},
|
context={"level": "high"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -853,7 +924,11 @@ async def test_engine_process_condition_gt(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond GT WF",
|
"name": "Cond GT WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -863,9 +938,13 @@ async def test_engine_process_condition_gt(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"count": 10},
|
context={"count": 10},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -888,7 +967,11 @@ async def test_engine_process_condition_lt(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond LT WF",
|
"name": "Cond LT WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -898,9 +981,13 @@ async def test_engine_process_condition_lt(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"count": 50},
|
context={"count": 50},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -923,7 +1010,11 @@ async def test_engine_process_condition_contains(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond Contains WF",
|
"name": "Cond Contains WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -933,9 +1024,13 @@ async def test_engine_process_condition_contains(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"tags": "vip,premium"},
|
context={"tags": "vip,premium"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -958,7 +1053,11 @@ async def test_engine_process_condition_contains_list(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond List WF",
|
"name": "Cond List WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -968,9 +1067,13 @@ async def test_engine_process_condition_contains_list(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"roles": ["admin", "editor"]},
|
context={"roles": ["admin", "editor"]},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -993,7 +1096,11 @@ async def test_engine_process_condition_no_match_advances(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond No Branch WF",
|
"name": "Cond No Branch WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -1003,9 +1110,13 @@ async def test_engine_process_condition_no_match_advances(db_session):
|
|||||||
},
|
},
|
||||||
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
{"name": "Next", "type": "action", "config": {"action_type": "noop"}},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"status": "inactive"},
|
context={"status": "inactive"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1028,7 +1139,11 @@ async def test_engine_process_condition_completes_last_step(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Cond Last WF",
|
"name": "Cond Last WF",
|
||||||
"steps": [
|
"steps": [
|
||||||
{
|
{
|
||||||
@@ -1037,9 +1152,13 @@ async def test_engine_process_condition_completes_last_step(db_session):
|
|||||||
"config": {"field": "status", "operator": "eq", "value": "active"},
|
"config": {"field": "status", "operator": "eq", "value": "active"},
|
||||||
},
|
},
|
||||||
],
|
],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"],
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
context={"status": "active"},
|
context={"status": "active"},
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1062,10 +1181,15 @@ async def test_engine_unknown_step_type(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Unknown Type WF",
|
"name": "Unknown Type WF",
|
||||||
"steps": [{"name": "Bad", "type": "unknown_type", "config": {}}],
|
"steps": [{"name": "Bad", "type": "unknown_type", "config": {}}],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1082,19 +1206,25 @@ async def test_engine_unknown_step_type(db_session):
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_engine_workflow_not_found(db_session):
|
async def test_engine_workflow_not_found(db_session):
|
||||||
"""Engine: process_step with non-existent workflow returns 404 error."""
|
"""Engine: process_step with non-existent workflow returns 404 error."""
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
from app.workflows.engine import WorkflowEngine
|
|
||||||
from app.models.workflow import WorkflowInstance
|
from app.models.workflow import WorkflowInstance
|
||||||
|
from app.workflows.engine import WorkflowEngine
|
||||||
|
|
||||||
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
|
||||||
|
|
||||||
# Create a real workflow and instance
|
# Create a real workflow and instance
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Temp WF for Not Found Test",
|
"name": "Temp WF for Not Found Test",
|
||||||
"steps": [{"name": "Step", "type": "action", "config": {"action_type": "noop"}}],
|
"steps": [{"name": "Step", "type": "action", "config": {"action_type": "noop"}}],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1130,10 +1260,15 @@ async def test_engine_step_index_beyond_steps_completes(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Beyond Steps WF",
|
"name": "Beyond Steps WF",
|
||||||
"steps": [{"name": "Step", "type": "action", "config": {"action_type": "noop"}}],
|
"steps": [{"name": "Step", "type": "action", "config": {"action_type": "noop"}}],
|
||||||
})
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1157,14 +1292,21 @@ async def test_engine_handle_event(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
await create_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
"name": "Event Handler WF",
|
"name": "Event Handler WF",
|
||||||
"trigger_event": "company.created",
|
"trigger_event": "company.created",
|
||||||
"steps": VALID_STEPS,
|
"steps": VALID_STEPS,
|
||||||
})
|
},
|
||||||
|
)
|
||||||
|
|
||||||
instances = await handle_event(
|
instances = await handle_event(
|
||||||
db_session, tenant_id, "company.created",
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
"company.created",
|
||||||
{"user_id": str(admin_id), "company_id": "test"},
|
{"user_id": str(admin_id), "company_id": "test"},
|
||||||
)
|
)
|
||||||
assert len(instances) >= 1
|
assert len(instances) >= 1
|
||||||
@@ -1180,7 +1322,9 @@ async def test_engine_handle_event_no_match(db_session):
|
|||||||
tenant_id = seed["tenant_a"].id
|
tenant_id = seed["tenant_a"].id
|
||||||
|
|
||||||
instances = await handle_event(
|
instances = await handle_event(
|
||||||
db_session, tenant_id, "nonexistent.event",
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
"nonexistent.event",
|
||||||
{"user_id": str(seed["admin_a"].id)},
|
{"user_id": str(seed["admin_a"].id)},
|
||||||
)
|
)
|
||||||
assert instances == []
|
assert instances == []
|
||||||
@@ -1189,8 +1333,9 @@ async def test_engine_handle_event_no_match(db_session):
|
|||||||
def test_engine_register_workflow_event_handlers():
|
def test_engine_register_workflow_event_handlers():
|
||||||
"""Engine: register_workflow_event_handlers subscribes to event bus."""
|
"""Engine: register_workflow_event_handlers subscribes to event bus."""
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from app.workflows.engine import register_workflow_event_handlers
|
|
||||||
from app.core.event_bus import get_event_bus
|
from app.core.event_bus import get_event_bus
|
||||||
|
from app.workflows.engine import register_workflow_event_handlers
|
||||||
|
|
||||||
event_bus = get_event_bus()
|
event_bus = get_event_bus()
|
||||||
# Reset handlers to a fresh defaultdict for test isolation
|
# Reset handlers to a fresh defaultdict for test isolation
|
||||||
@@ -1206,6 +1351,7 @@ def test_engine_register_workflow_event_handlers():
|
|||||||
|
|
||||||
# ─── Route Error Path Tests ───
|
# ─── Route Error Path Tests ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_route_update_workflow_not_found(client: AsyncClient, db_session):
|
async def test_route_update_workflow_not_found(client: AsyncClient, db_session):
|
||||||
"""Route: PATCH non-existent workflow returns 404."""
|
"""Route: PATCH non-existent workflow returns 404."""
|
||||||
@@ -1234,11 +1380,16 @@ async def test_route_update_workflow_rbac_blocked(client: AsyncClient, db_sessio
|
|||||||
wf_id = create_resp.json()["id"]
|
wf_id = create_resp.json()["id"]
|
||||||
|
|
||||||
# Login as viewer
|
# Login as viewer
|
||||||
from httpx import ASGITransport, AsyncClient as AC
|
from httpx import ASGITransport
|
||||||
|
from httpx import AsyncClient as AC # noqa: N817
|
||||||
|
|
||||||
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 viewer_client:
|
async with AC(
|
||||||
|
transport=ASGITransport(app=app_instance), base_url="http://test"
|
||||||
|
) as viewer_client:
|
||||||
await login_client(viewer_client, "viewer@tenanta.com")
|
await login_client(viewer_client, "viewer@tenanta.com")
|
||||||
resp = await viewer_client.patch(
|
resp = await viewer_client.patch(
|
||||||
f"/api/v1/workflows/{wf_id}",
|
f"/api/v1/workflows/{wf_id}",
|
||||||
@@ -1271,11 +1422,16 @@ async def test_route_delete_workflow_rbac_blocked(client: AsyncClient, db_sessio
|
|||||||
)
|
)
|
||||||
wf_id = create_resp.json()["id"]
|
wf_id = create_resp.json()["id"]
|
||||||
|
|
||||||
from httpx import ASGITransport, AsyncClient as AC
|
from httpx import ASGITransport
|
||||||
|
from httpx import AsyncClient as AC # noqa: N817
|
||||||
|
|
||||||
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 viewer_client:
|
async with AC(
|
||||||
|
transport=ASGITransport(app=app_instance), base_url="http://test"
|
||||||
|
) as viewer_client:
|
||||||
await login_client(viewer_client, "viewer@tenanta.com")
|
await login_client(viewer_client, "viewer@tenanta.com")
|
||||||
resp = await viewer_client.delete(f"/api/v1/workflows/{wf_id}", headers=ORIGIN_HEADER)
|
resp = await viewer_client.delete(f"/api/v1/workflows/{wf_id}", headers=ORIGIN_HEADER)
|
||||||
assert resp.status_code == 403
|
assert resp.status_code == 403
|
||||||
@@ -1330,7 +1486,10 @@ async def test_route_advance_completed_instance_returns_400(client: AsyncClient,
|
|||||||
|
|
||||||
create_resp = await client.post(
|
create_resp = await client.post(
|
||||||
"/api/v1/workflows",
|
"/api/v1/workflows",
|
||||||
json={"name": "Single Step", "steps": [{"name": "Only", "type": "action", "config": {"action_type": "noop"}}]},
|
json={
|
||||||
|
"name": "Single Step",
|
||||||
|
"steps": [{"name": "Only", "type": "action", "config": {"action_type": "noop"}}],
|
||||||
|
},
|
||||||
headers=ORIGIN_HEADER,
|
headers=ORIGIN_HEADER,
|
||||||
)
|
)
|
||||||
wf_id = create_resp.json()["id"]
|
wf_id = create_resp.json()["id"]
|
||||||
@@ -1395,6 +1554,7 @@ async def test_route_cancel_rejected_instance_returns_400(client: AsyncClient, d
|
|||||||
|
|
||||||
# ─── Service Edge Case Tests ───
|
# ─── Service Edge Case Tests ───
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_service_check_timeout_no_timeout_at(db_session):
|
async def test_service_check_timeout_no_timeout_at(db_session):
|
||||||
"""Service: check_timeout returns False when timeout_at is None."""
|
"""Service: check_timeout returns False when timeout_at is None."""
|
||||||
@@ -1402,9 +1562,15 @@ async def test_service_check_timeout_no_timeout_at(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
"name": "No Timeout WF", "steps": VALID_STEPS,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "No Timeout WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"])
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1423,11 +1589,21 @@ async def test_service_check_timeout_completed_status(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
"name": "Completed Timeout WF", "steps": VALID_STEPS,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "Completed Timeout WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(
|
inst = await create_instance(
|
||||||
db_session, tenant_id, admin_id, wf["id"], timeout_hours=1,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
|
timeout_hours=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1435,7 +1611,7 @@ async def test_service_check_timeout_completed_status(db_session):
|
|||||||
)
|
)
|
||||||
instance = result.scalar_one()
|
instance = result.scalar_one()
|
||||||
instance.status = "completed"
|
instance.status = "completed"
|
||||||
instance.timeout_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
instance.timeout_at = datetime.now(UTC) - timedelta(hours=1)
|
||||||
await db_session.flush()
|
await db_session.flush()
|
||||||
|
|
||||||
is_timed_out = await check_timeout(instance)
|
is_timed_out = await check_timeout(instance)
|
||||||
@@ -1449,9 +1625,15 @@ async def test_service_auto_reject_no_initiated_by(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
"name": "No Initiator WF", "steps": VALID_STEPS,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "No Initiator WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
},
|
||||||
|
)
|
||||||
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"], timeout_hours=1)
|
inst = await create_instance(db_session, tenant_id, admin_id, wf["id"], timeout_hours=1)
|
||||||
|
|
||||||
result = await db_session.execute(
|
result = await db_session.execute(
|
||||||
@@ -1459,7 +1641,7 @@ async def test_service_auto_reject_no_initiated_by(db_session):
|
|||||||
)
|
)
|
||||||
instance = result.scalar_one()
|
instance = result.scalar_one()
|
||||||
instance.initiated_by = None
|
instance.initiated_by = None
|
||||||
instance.timeout_at = datetime.now(timezone.utc) - timedelta(hours=1)
|
instance.timeout_at = datetime.now(UTC) - timedelta(hours=1)
|
||||||
await db_session.flush()
|
await db_session.flush()
|
||||||
|
|
||||||
reject_result = await auto_reject_timeout(db_session, tenant_id, instance)
|
reject_result = await auto_reject_timeout(db_session, tenant_id, instance)
|
||||||
@@ -1484,7 +1666,9 @@ async def test_service_start_instance_no_match(db_session):
|
|||||||
admin_id = seed["admin_a"].id
|
admin_id = seed["admin_a"].id
|
||||||
|
|
||||||
instances = await start_instance_for_event(
|
instances = await start_instance_for_event(
|
||||||
db_session, tenant_id, admin_id,
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
"nonexistent.event",
|
"nonexistent.event",
|
||||||
context={},
|
context={},
|
||||||
)
|
)
|
||||||
@@ -1499,13 +1683,27 @@ async def test_service_list_workflows_with_active_filter(db_session):
|
|||||||
admin_id = seed["admin_a"].id
|
admin_id = seed["admin_a"].id
|
||||||
|
|
||||||
# Create active workflow
|
# Create active workflow
|
||||||
await create_workflow(db_session, tenant_id, admin_id, {
|
await create_workflow(
|
||||||
"name": "Active WF", "steps": VALID_STEPS, "is_active": True,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "Active WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
"is_active": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
# Create inactive workflow
|
# Create inactive workflow
|
||||||
await create_workflow(db_session, tenant_id, admin_id, {
|
await create_workflow(
|
||||||
"name": "Inactive WF", "steps": VALID_STEPS, "is_active": False,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "Inactive WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
"is_active": False,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
result = await list_workflows(db_session, tenant_id, is_active=True)
|
result = await list_workflows(db_session, tenant_id, is_active=True)
|
||||||
assert result["total"] == 1
|
assert result["total"] == 1
|
||||||
@@ -1519,15 +1717,27 @@ async def test_service_update_workflow_with_steps(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
|
||||||
|
|
||||||
wf = await create_workflow(db_session, tenant_id, admin_id, {
|
wf = await create_workflow(
|
||||||
"name": "Steps Update WF", "steps": VALID_STEPS,
|
db_session,
|
||||||
})
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
{
|
||||||
|
"name": "Steps Update WF",
|
||||||
|
"steps": VALID_STEPS,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
new_steps = [{"name": "New Step", "type": "action", "config": {"action_type": "noop"}}]
|
new_steps = [{"name": "New Step", "type": "action", "config": {"action_type": "noop"}}]
|
||||||
result = await update_workflow(db_session, tenant_id, admin_id, wf["id"], {
|
result = await update_workflow(
|
||||||
|
db_session,
|
||||||
|
tenant_id,
|
||||||
|
admin_id,
|
||||||
|
wf["id"],
|
||||||
|
{
|
||||||
"steps": new_steps,
|
"steps": new_steps,
|
||||||
"is_active": False,
|
"is_active": False,
|
||||||
})
|
},
|
||||||
|
)
|
||||||
assert result is not None
|
assert result is not None
|
||||||
assert len(result["steps"]) == 1
|
assert len(result["steps"]) == 1
|
||||||
assert result["steps"][0]["name"] == "New Step"
|
assert result["steps"][0]["name"] == "New Step"
|
||||||
@@ -1541,7 +1751,9 @@ async def test_service_update_workflow_not_found(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 update_workflow(db_session, tenant_id, admin_id, str(uuid.uuid4()), {"name": "X"})
|
result = await update_workflow(
|
||||||
|
db_session, tenant_id, admin_id, str(uuid.uuid4()), {"name": "X"}
|
||||||
|
)
|
||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
@@ -1559,7 +1771,6 @@ async def test_service_delete_workflow_not_found(db_session):
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_service_get_instance_not_found(db_session):
|
async def test_service_get_instance_not_found(db_session):
|
||||||
"""Service: get_instance returns None for non-existent instance."""
|
"""Service: get_instance returns None for non-existent instance."""
|
||||||
from app.services.workflow_service import get_instance
|
|
||||||
|
|
||||||
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
|
||||||
|
|||||||
Reference in New Issue
Block a user