Merge pull request #554 from christianlouis/copilot/add-api-subscription-quota-enforcement

feat(integrations): add subscription quota enforcement, connection test, and quota endpoint
This commit is contained in:
Christian Krakau-Louis
2026-03-08 18:11:15 +01:00
committed by GitHub
3 changed files with 970 additions and 1 deletions
+344
View File
@@ -7,6 +7,15 @@ storage destination (e.g. S3, Dropbox, Google Drive) configured by a user.
Sensitive credentials are encrypted at rest using Fernet symmetric encryption Sensitive credentials are encrypted at rest using Fernet symmetric encryption
(keyed from ``SESSION_SECRET``) via :mod:`app.utils.encryption`. Credential (keyed from ``SESSION_SECRET``) via :mod:`app.utils.encryption`. Credential
values are **never** returned in API responses. values are **never** returned in API responses.
Subscription quota enforcement
------------------------------
On creation, the endpoint checks the user's subscription tier limits:
* **Destinations** — ``max_storage_destinations`` from the plan.
* **Sources (IMAP)** — ``max_mailboxes`` from the plan.
Exceeding the quota returns HTTP 403 with an actionable error message.
""" """
import json import json
@@ -20,6 +29,7 @@ from sqlalchemy.orm import Session
from app.database import get_db from app.database import get_db
from app.models import IntegrationDirection, IntegrationType, UserIntegration from app.models import IntegrationDirection, IntegrationType, UserIntegration
from app.utils.encryption import decrypt_value, encrypt_value from app.utils.encryption import decrypt_value, encrypt_value
from app.utils.subscription import get_tier, get_user_tier_id
from app.utils.user_scope import get_current_owner_id from app.utils.user_scope import get_current_owner_id
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -42,6 +52,119 @@ def _get_owner_id(request: Request) -> str:
CurrentOwner = Annotated[str, Depends(_get_owner_id)] CurrentOwner = Annotated[str, Depends(_get_owner_id)]
# ---------------------------------------------------------------------------
# Quota helpers
# ---------------------------------------------------------------------------
_FREE_TIER_ID = "free"
# Source types that consume the mailbox quota
_MAILBOX_SOURCE_TYPES = {IntegrationType.IMAP}
def _get_max_destinations(tier: dict[str, Any]) -> int | None:
"""Return the maximum number of storage destinations allowed by *tier*.
Returns:
``None`` — unlimited (paid tiers with ``max_storage_destinations == 0``)
positive — the configured limit
"""
tier_id: str = tier.get("id", _FREE_TIER_ID)
max_dest: int = tier.get("max_storage_destinations", 0)
# Free tier: the value itself is the limit (e.g. 1)
if tier_id == _FREE_TIER_ID:
return max_dest if max_dest > 0 else 1 # safe default
# Paid tiers: 0 means unlimited
if max_dest == 0:
return None
return max_dest
def _get_max_sources(tier: dict[str, Any]) -> int | None:
"""Return the maximum number of IMAP source integrations allowed by *tier*.
Returns:
``None`` — unlimited (paid tiers with ``max_mailboxes == 0``)
``0`` — no mailboxes allowed (free tier)
positive — the configured limit
"""
tier_id: str = tier.get("id", _FREE_TIER_ID)
max_mb: int = tier.get("max_mailboxes", 0)
# Free tier: 0 means "no access" (not "unlimited")
if tier_id == _FREE_TIER_ID:
return 0
# Paid tiers: 0 means unlimited
if max_mb == 0:
return None
return max_mb
def _check_quota(db: Session, owner_id: str, direction: str, integration_type: str) -> None:
"""Raise 403 if the user has reached their integration quota.
Quota rules:
* DESTINATION integrations are limited by ``max_storage_destinations``.
* SOURCE integrations of type IMAP are limited by ``max_mailboxes``.
* Other SOURCE types (WATCH_FOLDER, WEBHOOK) are not quota-limited yet.
"""
tier_id = get_user_tier_id(db, owner_id)
tier = get_tier(tier_id, db)
if direction == IntegrationDirection.DESTINATION:
max_dest = _get_max_destinations(tier)
if max_dest is not None:
current_count = (
db.query(UserIntegration)
.filter(
UserIntegration.owner_id == owner_id,
UserIntegration.direction == IntegrationDirection.DESTINATION,
)
.count()
)
if current_count >= max_dest:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=(
f"You have reached your plan limit of {max_dest} storage destination(s). "
"Please remove an existing destination or upgrade your plan."
),
)
elif direction == IntegrationDirection.SOURCE and integration_type in _MAILBOX_SOURCE_TYPES:
max_src = _get_max_sources(tier)
if max_src == 0:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="Your current plan does not include email ingestion. Upgrade to a paid plan to add IMAP sources.",
)
if max_src is not None:
current_count = (
db.query(UserIntegration)
.filter(
UserIntegration.owner_id == owner_id,
UserIntegration.direction == IntegrationDirection.SOURCE,
UserIntegration.integration_type.in_(list(_MAILBOX_SOURCE_TYPES)),
)
.count()
)
if current_count >= max_src:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=(
f"You have reached your plan limit of {max_src} IMAP source(s). "
"Please remove an existing source or upgrade your plan."
),
)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Pydantic schemas # Pydantic schemas
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -72,6 +195,14 @@ class IntegrationUpdate(BaseModel):
is_active: bool | None = None is_active: bool | None = None
class IntegrationTestRequest(BaseModel):
"""Schema for testing an integration connection without saving it."""
integration_type: str = Field(..., description="Integration type (e.g. 'IMAP', 'S3', 'DROPBOX')")
config: dict[str, Any] | None = Field(default=None, description="Non-sensitive configuration")
credentials: dict[str, Any] | None = Field(default=None, description="Credentials for the connection test")
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Validation helpers # Validation helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -199,10 +330,15 @@ def create_integration(
``credentials`` are encrypted at rest using Fernet symmetric encryption ``credentials`` are encrypted at rest using Fernet symmetric encryption
before being persisted and are **never** returned in API responses. before being persisted and are **never** returned in API responses.
Quota is enforced against the user's subscription plan before the
integration is persisted.
""" """
_validate_direction(body.direction) _validate_direction(body.direction)
_validate_integration_type(body.integration_type) _validate_integration_type(body.integration_type)
_check_quota(db, owner_id, body.direction, body.integration_type)
integration = UserIntegration( integration = UserIntegration(
owner_id=owner_id, owner_id=owner_id,
direction=body.direction, direction=body.direction,
@@ -342,3 +478,211 @@ def get_integration_credentials(
credentials = _decode_credentials(integration.credentials) credentials = _decode_credentials(integration.credentials)
return {"credentials": credentials or {}} return {"credentials": credentials or {}}
# ---------------------------------------------------------------------------
# Connection test helpers
# ---------------------------------------------------------------------------
def _test_imap_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
"""Test an IMAP connection using the provided config and credentials."""
import imaplib
cfg = config or {}
creds = credentials or {}
host = cfg.get("host", "")
port = int(cfg.get("port", 993))
username = cfg.get("username", "")
password = creds.get("password", "")
use_ssl = cfg.get("use_ssl", True)
if not host or not username or not password:
return {"success": False, "message": "Missing required fields: host, username, and password"}
try:
if use_ssl:
mail = imaplib.IMAP4_SSL(host, port)
else:
mail = imaplib.IMAP4(host, port)
mail.login(username, password)
mail.logout()
return {"success": True, "message": "IMAP connection successful"}
except OSError as exc:
logger.warning("IMAP network error for %s@%s: %s", username, host, exc)
return {"success": False, "message": "IMAP connection failed — check host, port, and network connectivity"}
except Exception as exc: # noqa: BLE001
logger.warning("IMAP error for %s@%s: %s", username, host, exc)
return {"success": False, "message": "IMAP authentication or connection failed"}
def _test_s3_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
"""Test an S3 connection by calling HeadBucket."""
try:
import boto3
from botocore.exceptions import BotoCoreError, ClientError
except ImportError:
return {"success": False, "message": "boto3 is not installed"}
cfg = config or {}
creds = credentials or {}
bucket = cfg.get("bucket", "")
region = cfg.get("region", "us-east-1")
if not bucket:
return {"success": False, "message": "Missing required field: bucket"}
try:
client = boto3.client(
"s3",
region_name=region,
aws_access_key_id=creds.get("access_key_id", ""),
aws_secret_access_key=creds.get("secret_access_key", ""),
endpoint_url=cfg.get("endpoint_url"),
)
client.head_bucket(Bucket=bucket)
return {"success": True, "message": f"S3 bucket '{bucket}' is accessible"}
except (BotoCoreError, ClientError) as exc:
logger.warning("S3 connection error for bucket '%s': %s", bucket, exc)
return {"success": False, "message": "S3 connection failed — check bucket name, region, and credentials"}
except Exception as exc: # noqa: BLE001
logger.warning("S3 unexpected error for bucket '%s': %s", bucket, exc)
return {"success": False, "message": "S3 connection failed"}
def _test_webdav_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
"""Test a WebDAV/Nextcloud connection by issuing an HTTP PROPFIND."""
import urllib.request
cfg = config or {}
creds = credentials or {}
url = cfg.get("url", "")
username = creds.get("username", "")
password = creds.get("password", "")
if not url:
return {"success": False, "message": "Missing required field: url"}
# Only allow http/https to prevent file:// or other custom scheme attacks
import ipaddress
from urllib.parse import urlparse
parsed = urlparse(url)
if parsed.scheme not in ("http", "https"):
return {"success": False, "message": "URL must use http or https scheme"}
# Block requests to private/internal IPs to prevent SSRF
hostname = parsed.hostname or ""
if hostname:
try:
addr = ipaddress.ip_address(hostname)
if addr.is_private or addr.is_loopback or addr.is_link_local:
return {"success": False, "message": "URLs pointing to internal or private networks are not allowed"}
except ValueError:
# Hostname is not an IP literal — allow DNS names through
if hostname in ("localhost", "localhost.localdomain"):
return {"success": False, "message": "URLs pointing to localhost are not allowed"}
try:
import base64
req = urllib.request.Request(url, method="PROPFIND") # noqa: S310
if username and password:
token = base64.b64encode(f"{username}:{password}".encode()).decode()
req.add_header("Authorization", f"Basic {token}")
req.add_header("Depth", "0")
with urllib.request.urlopen(req, timeout=10) as resp: # noqa: S310
if resp.status < 400:
return {"success": True, "message": "WebDAV connection successful"}
return {"success": False, "message": f"WebDAV returned HTTP {resp.status}"}
except Exception as exc: # noqa: BLE001
logger.warning("WebDAV connection error for %s: %s", hostname, exc)
return {"success": False, "message": "WebDAV connection failed — check URL and credentials"}
_CONNECTION_TESTERS: dict[str, Any] = {
IntegrationType.IMAP: _test_imap_connection,
IntegrationType.S3: _test_s3_connection,
IntegrationType.WEBDAV: _test_webdav_connection,
IntegrationType.NEXTCLOUD: _test_webdav_connection,
}
# ---------------------------------------------------------------------------
# Test & quota endpoints
# ---------------------------------------------------------------------------
@router.post("/test", summary="Test an integration connection without saving")
def test_integration_connection(
request: Request,
body: IntegrationTestRequest,
owner_id: CurrentOwner,
) -> dict[str, Any]:
"""Test integration credentials without persisting anything.
Useful for the "Test connection" button in the UI before the user saves
a new integration. Returns ``{"success": bool, "message": str}``.
"""
_validate_integration_type(body.integration_type)
tester = _CONNECTION_TESTERS.get(body.integration_type)
if tester is None:
return {
"success": False,
"message": f"Connection testing is not yet supported for '{body.integration_type}'. "
"The integration can still be saved and will be validated on first use.",
}
return tester(body.config, body.credentials)
@router.get("/quota/", summary="Get integration quota information for the current user")
def get_integration_quota(
request: Request,
db: DbSession,
owner_id: CurrentOwner,
) -> dict[str, Any]:
"""Return the user's current integration usage vs. their plan quota.
Includes separate counts for destinations and IMAP sources.
"""
tier_id = get_user_tier_id(db, owner_id)
tier = get_tier(tier_id, db)
max_dest = _get_max_destinations(tier)
max_src = _get_max_sources(tier)
dest_count = (
db.query(UserIntegration)
.filter(
UserIntegration.owner_id == owner_id,
UserIntegration.direction == IntegrationDirection.DESTINATION,
)
.count()
)
src_count = (
db.query(UserIntegration)
.filter(
UserIntegration.owner_id == owner_id,
UserIntegration.direction == IntegrationDirection.SOURCE,
UserIntegration.integration_type.in_(list(_MAILBOX_SOURCE_TYPES)),
)
.count()
)
return {
"tier_id": tier_id,
"tier_name": tier.get("name", tier_id),
"destinations": {
"current_count": dest_count,
"max_allowed": max_dest,
"can_add": max_dest is None or dest_count < max_dest,
},
"sources": {
"current_count": src_count,
"max_allowed": max_src,
"can_add": max_src is None or (max_src > 0 and src_count < max_src),
},
}
+129
View File
@@ -1218,6 +1218,135 @@ Send a processed file to Google Drive.
} }
``` ```
## Integrations
Manage per-user integrations (sources and destinations). All endpoints require authentication and are scoped to the current user's integrations. Subscription-tier quota enforcement is applied on creation.
### Quota Enforcement
When creating an integration, the API checks the user's subscription tier:
| Tier | Storage Destinations | IMAP Sources |
|------|---------------------|--------------|
| **Free** | 1 | 0 |
| **Starter** | 2 | 1 |
| **Professional** | 5 | 3 |
| **Power** | 10 | Unlimited |
Exceeding a quota returns HTTP 403 with a descriptive error message.
### GET /api/integrations/
List all integrations for the current user. Supports optional query-string filters.
**Query Parameters:**
| Parameter | Type | Description |
|-----------|------|-------------|
| `direction` | string | Filter by `SOURCE` or `DESTINATION` |
| `integration_type` | string | Filter by type (e.g. `IMAP`, `S3`, `DROPBOX`) |
**Response (200):**
```json
[
{
"id": 1,
"owner_id": "user@example.com",
"direction": "DESTINATION",
"integration_type": "S3",
"name": "Archive Bucket",
"config": {"bucket": "my-bucket", "region": "us-east-1"},
"has_credentials": true,
"is_active": true,
"last_used_at": null,
"last_error": null,
"created_at": "2025-01-01T00:00:00",
"updated_at": "2025-01-01T00:00:00"
}
]
```
### POST /api/integrations/
Create a new integration. Quota is enforced before creation.
**Request:**
```json
{
"direction": "DESTINATION",
"integration_type": "S3",
"name": "Archive Bucket",
"config": {"bucket": "my-bucket", "region": "us-east-1"},
"credentials": {"access_key_id": "AKIA...", "secret_access_key": "..."},
"is_active": true
}
```
**Response (201):** The created integration (same shape as list response).
**Response (403):** Quota exceeded.
```json
{
"detail": "You have reached your plan limit of 1 storage destination(s). Please remove an existing destination or upgrade your plan."
}
```
### PUT /api/integrations/{id}
Update an existing integration. Only provided fields are changed.
### DELETE /api/integrations/{id}
Delete an integration permanently. Returns 204 on success.
### POST /api/integrations/test
Test an integration connection without saving. Useful for "Test connection" UI buttons.
**Request:**
```json
{
"integration_type": "IMAP",
"config": {"host": "imap.gmail.com", "port": 993, "username": "user@example.com", "use_ssl": true},
"credentials": {"password": "app-password"}
}
```
**Response (200):**
```json
{"success": true, "message": "IMAP connection successful"}
```
Supported connection tests: `IMAP`, `S3`, `WEBDAV`, `NEXTCLOUD`. Other types return a message that testing is not yet supported.
### GET /api/integrations/quota/
Get the current user's integration quota usage.
**Response (200):**
```json
{
"tier_id": "starter",
"tier_name": "Starter",
"destinations": {
"current_count": 1,
"max_allowed": 2,
"can_add": true
},
"sources": {
"current_count": 0,
"max_allowed": 1,
"can_add": true
}
}
```
## Webhooks ## Webhooks
Manage webhook configurations for notifying external systems when document events occur. All webhook endpoints require admin access. Manage webhook configurations for notifying external systems when document events occur. All webhook endpoints require admin access.
+497 -1
View File
@@ -7,7 +7,7 @@ from sqlalchemy.orm import sessionmaker
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from app.database import Base, get_db from app.database import Base, get_db
from app.models import UserIntegration from app.models import SubscriptionPlan, UserIntegration, UserProfile
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Test data constants # Test data constants
@@ -43,6 +43,37 @@ _S3_DESTINATION = {
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _make_profile(session, owner: str = _OWNER, tier: str = "business") -> UserProfile:
"""Create a UserProfile row for the given user and tier."""
profile = UserProfile(user_id=owner, subscription_tier=tier)
session.add(profile)
session.commit()
session.refresh(profile)
return profile
def _make_plan(
session,
tier: str = "business",
max_storage_destinations: int = 10,
max_mailboxes: int = 0,
) -> SubscriptionPlan:
"""Create a SubscriptionPlan row."""
plan = SubscriptionPlan(
plan_id=tier,
name=tier.title(),
price_monthly=7.99,
price_yearly=76.99,
max_storage_destinations=max_storage_destinations,
max_mailboxes=max_mailboxes,
is_active=True,
)
session.add(plan)
session.commit()
session.refresh(plan)
return plan
@pytest.fixture() @pytest.fixture()
def int_engine(): def int_engine():
"""In-memory SQLite engine for integration tests.""" """In-memory SQLite engine for integration tests."""
@@ -65,11 +96,41 @@ def int_session(int_engine):
session.close() session.close()
def _seed_default_plan(engine, owner_id: str = _OWNER) -> None:
"""Seed a generous (business-tier) plan and profile for the given user.
Called automatically by ``_make_client`` so that existing CRUD tests keep
working after quota enforcement was added.
"""
Session = sessionmaker(bind=engine)
session = Session()
try:
if not session.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == "business").first():
session.add(
SubscriptionPlan(
plan_id="business",
name="Power",
price_monthly=7.99,
price_yearly=76.99,
max_storage_destinations=10,
max_mailboxes=0, # 0 = unlimited for paid tiers
is_active=True,
)
)
if not session.query(UserProfile).filter(UserProfile.user_id == owner_id).first():
session.add(UserProfile(user_id=owner_id, subscription_tier="business"))
session.commit()
finally:
session.close()
def _make_client(int_engine, owner_id: str = _OWNER): def _make_client(int_engine, owner_id: str = _OWNER):
"""Return a TestClient with *owner_id* injected as the authenticated user.""" """Return a TestClient with *owner_id* injected as the authenticated user."""
from app.api.integrations import _get_owner_id from app.api.integrations import _get_owner_id
from app.main import app from app.main import app
_seed_default_plan(int_engine, owner_id)
def override_db(): def override_db():
Session = sessionmaker(bind=int_engine) Session = sessionmaker(bind=int_engine)
session = Session() session = Session()
@@ -548,3 +609,438 @@ class TestImapPasswordEncryption:
verify_session.close() verify_session.close()
finally: finally:
app.dependency_overrides.clear() app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# Quota enforcement tests
# ---------------------------------------------------------------------------
@pytest.mark.integration
class TestQuotaEnforcementDestinations:
"""Tests for destination quota enforcement on POST /api/integrations/."""
def test_create_destination_blocked_at_limit(self, int_engine, int_session):
"""Users at the destination quota limit receive a 403."""
_make_profile(int_session, tier="starter")
_make_plan(int_session, tier="starter", max_storage_destinations=1, max_mailboxes=1)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
# First destination should succeed
resp1 = client.post("/api/integrations/", json=_S3_DESTINATION)
assert resp1.status_code == 201
# Second destination should be blocked
second = dict(_S3_DESTINATION, name="Second Bucket")
resp2 = client.post("/api/integrations/", json=second)
assert resp2.status_code == 403
assert "limit" in resp2.json()["detail"].lower()
finally:
app.dependency_overrides.clear()
def test_create_destination_allowed_under_limit(self, int_engine, int_session):
"""Users under the destination quota can create integrations."""
_make_profile(int_session, tier="professional")
_make_plan(int_session, tier="professional", max_storage_destinations=5, max_mailboxes=3)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp = client.post("/api/integrations/", json=_S3_DESTINATION)
assert resp.status_code == 201
finally:
app.dependency_overrides.clear()
def test_free_tier_allows_one_destination(self, int_engine, int_session):
"""Free tier allows exactly 1 destination."""
_make_profile(int_session, tier="free")
_make_plan(int_session, tier="free", max_storage_destinations=1, max_mailboxes=0)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp1 = client.post("/api/integrations/", json=_S3_DESTINATION)
assert resp1.status_code == 201
second = dict(_S3_DESTINATION, name="Second")
resp2 = client.post("/api/integrations/", json=second)
assert resp2.status_code == 403
finally:
app.dependency_overrides.clear()
@pytest.mark.integration
class TestQuotaEnforcementSources:
"""Tests for IMAP source quota enforcement on POST /api/integrations/."""
def test_create_imap_source_blocked_on_free_tier(self, int_engine, int_session):
"""Free-tier users cannot add IMAP source integrations."""
_make_profile(int_session, tier="free")
_make_plan(int_session, tier="free", max_storage_destinations=1, max_mailboxes=0)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp = client.post("/api/integrations/", json=_IMAP_SOURCE)
assert resp.status_code == 403
assert "plan" in resp.json()["detail"].lower()
finally:
app.dependency_overrides.clear()
def test_create_imap_source_blocked_at_limit(self, int_engine, int_session):
"""Starter-tier users with 1 IMAP source cannot add a second."""
_make_profile(int_session, tier="starter")
_make_plan(int_session, tier="starter", max_storage_destinations=2, max_mailboxes=1)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp1 = client.post("/api/integrations/", json=_IMAP_SOURCE)
assert resp1.status_code == 201
second = dict(_IMAP_SOURCE, name="Second Mailbox")
resp2 = client.post("/api/integrations/", json=second)
assert resp2.status_code == 403
finally:
app.dependency_overrides.clear()
def test_create_imap_source_unlimited_on_power_tier(self, int_engine, int_session):
"""Power-tier users can add multiple IMAP sources (unlimited)."""
_make_profile(int_session, tier="business")
_make_plan(int_session, tier="business", max_storage_destinations=10, max_mailboxes=0)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp1 = client.post("/api/integrations/", json=_IMAP_SOURCE)
resp2 = client.post("/api/integrations/", json=dict(_IMAP_SOURCE, name="Second"))
resp3 = client.post("/api/integrations/", json=dict(_IMAP_SOURCE, name="Third"))
assert resp1.status_code == 201
assert resp2.status_code == 201
assert resp3.status_code == 201
finally:
app.dependency_overrides.clear()
def test_watch_folder_source_not_quota_limited(self, int_engine, int_session):
"""WATCH_FOLDER sources are not subject to mailbox quota limits."""
_make_profile(int_session, tier="free")
_make_plan(int_session, tier="free", max_storage_destinations=1, max_mailboxes=0)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
payload = {
"direction": "SOURCE",
"integration_type": "WATCH_FOLDER",
"name": "My Folder",
"config": {"path": "/tmp/watch"},
"is_active": True,
}
resp = client.post("/api/integrations/", json=payload)
assert resp.status_code == 201
finally:
app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# Quota helpers unit tests
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestQuotaHelpers:
"""Unit tests for the quota helper functions."""
def test_get_max_destinations_free_tier(self):
from app.api.integrations import _get_max_destinations
assert _get_max_destinations({"id": "free", "max_storage_destinations": 1}) == 1
def test_get_max_destinations_paid_explicit(self):
from app.api.integrations import _get_max_destinations
assert _get_max_destinations({"id": "starter", "max_storage_destinations": 2}) == 2
def test_get_max_destinations_paid_unlimited(self):
from app.api.integrations import _get_max_destinations
assert _get_max_destinations({"id": "business", "max_storage_destinations": 0}) is None
def test_get_max_sources_free_tier(self):
from app.api.integrations import _get_max_sources
assert _get_max_sources({"id": "free", "max_mailboxes": 0}) == 0
def test_get_max_sources_paid_explicit(self):
from app.api.integrations import _get_max_sources
assert _get_max_sources({"id": "starter", "max_mailboxes": 1}) == 1
def test_get_max_sources_paid_unlimited(self):
from app.api.integrations import _get_max_sources
assert _get_max_sources({"id": "business", "max_mailboxes": 0}) is None
# ---------------------------------------------------------------------------
# Connection test endpoint tests
# ---------------------------------------------------------------------------
@pytest.mark.integration
class TestConnectionTestEndpoint:
"""Tests for POST /api/integrations/test."""
def test_test_unsupported_type(self, int_client):
"""Unsupported integration types return a helpful non-error message."""
payload = {
"integration_type": "DROPBOX",
"config": {},
"credentials": {"token": "abc"},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "not yet supported" in data["message"]
def test_test_invalid_type_returns_400(self, int_client):
"""Invalid integration_type returns 400."""
payload = {
"integration_type": "INVALID",
"config": {},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 400
def test_test_imap_missing_fields(self, int_client):
"""IMAP test with missing fields returns failure."""
payload = {
"integration_type": "IMAP",
"config": {"host": ""},
"credentials": {},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "Missing" in data["message"]
def test_test_s3_missing_bucket(self, int_client):
"""S3 test with missing bucket returns failure."""
payload = {
"integration_type": "S3",
"config": {},
"credentials": {"access_key_id": "AKIA", "secret_access_key": "secret"},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "bucket" in data["message"].lower()
def test_test_webdav_missing_url(self, int_client):
"""WebDAV test with missing URL returns failure."""
payload = {
"integration_type": "WEBDAV",
"config": {},
"credentials": {"username": "u", "password": "p"},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "url" in data["message"].lower()
def test_test_webdav_blocks_private_ip(self, int_client):
"""WebDAV test blocks requests to private/internal IPs (SSRF protection)."""
payload = {
"integration_type": "WEBDAV",
"config": {"url": "http://127.0.0.1/webdav"},
"credentials": {"username": "u", "password": "p"},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "internal" in data["message"].lower() or "private" in data["message"].lower()
def test_test_webdav_blocks_localhost(self, int_client):
"""WebDAV test blocks requests to localhost."""
payload = {
"integration_type": "WEBDAV",
"config": {"url": "http://localhost/webdav"},
"credentials": {},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "localhost" in data["message"].lower()
def test_test_webdav_blocks_file_scheme(self, int_client):
"""WebDAV test blocks file:// scheme."""
payload = {
"integration_type": "WEBDAV",
"config": {"url": "file:///etc/passwd"},
"credentials": {},
}
resp = int_client.post("/api/integrations/test", json=payload)
assert resp.status_code == 200
data = resp.json()
assert data["success"] is False
assert "scheme" in data["message"].lower()
# ---------------------------------------------------------------------------
# Quota endpoint tests
# ---------------------------------------------------------------------------
@pytest.mark.integration
class TestQuotaEndpoint:
"""Tests for GET /api/integrations/quota/."""
def test_quota_returns_tier_info(self, int_client):
"""Quota endpoint returns tier information and counts."""
resp = int_client.get("/api/integrations/quota/")
assert resp.status_code == 200
data = resp.json()
assert "tier_id" in data
assert "tier_name" in data
assert "destinations" in data
assert "sources" in data
assert "current_count" in data["destinations"]
assert "max_allowed" in data["destinations"]
assert "can_add" in data["destinations"]
def test_quota_reflects_created_integrations(self, int_client):
"""Quota counts update after creating integrations."""
int_client.post("/api/integrations/", json=_S3_DESTINATION)
resp = int_client.get("/api/integrations/quota/")
data = resp.json()
assert data["destinations"]["current_count"] == 1
def test_quota_free_tier(self, int_engine, int_session):
"""Free tier shows correct quota limits."""
_make_profile(int_session, tier="free")
_make_plan(int_session, tier="free", max_storage_destinations=1, max_mailboxes=0)
from app.api.integrations import _get_owner_id
from app.main import app
def override_db():
Session = sessionmaker(bind=int_engine)
session = Session()
try:
yield session
finally:
session.close()
app.dependency_overrides[get_db] = override_db
app.dependency_overrides[_get_owner_id] = lambda: _OWNER
try:
with TestClient(app, base_url="http://localhost", raise_server_exceptions=False) as client:
resp = client.get("/api/integrations/quota/")
data = resp.json()
assert data["tier_id"] == "free"
assert data["destinations"]["max_allowed"] == 1
assert data["destinations"]["can_add"] is True
assert data["sources"]["max_allowed"] == 0
assert data["sources"]["can_add"] is False
finally:
app.dependency_overrides.clear()