- Implement services/cloud_analysis_processing.py with full pipeline: queued -> downloading -> extracting -> classifying -> indexed - Stale guard: version key comparison fires before provider byte fetch (T-14-13) - Cache pin lifecycle: pin acquired before bytes used, released in finally block (T-14-12) - No provider mutation methods called: adapter only used via get_object (T-14-03) - No local Document row created: classification targets CloudItem directly - Cooperative cancellation check at job/item level before each stage - ProcessingResult dataclass with status, cache_entry_id, error_code fields - Add 4 contract tests: status_transitions, stale_detection, pin_release, no_mutations
823 lines
32 KiB
Python
823 lines
32 KiB
Python
"""
|
|
Cloud analysis processing service — Phase 14 Plan 05.
|
|
|
|
Implements the per-item processing pipeline:
|
|
1. Revalidate owner / connection / item from DB.
|
|
2. Check cooperative cancellation (job or item already cancelled).
|
|
3. Recompute version key from current CloudItem metadata; skip if stale.
|
|
4. Check cache — reuse non-evicted matching bytes or hydrate from provider.
|
|
5. Extract text from bytes (delegates to services.extractor).
|
|
6. Persist extracted_text onto CloudItem.
|
|
7. Classify topics through the AI provider (delegates to services.classifier
|
|
adapted for cloud items).
|
|
8. Update CloudItem.analysis_status and CloudAnalysisJobItem.status.
|
|
9. Update aggregate job counters.
|
|
10. Release cache pins in a finally block regardless of outcome.
|
|
|
|
Design invariants (enforced in tests and via FakeCloudAdapter):
|
|
- No provider MUTATION methods called (ANALYZE-07, T-14-03).
|
|
- No local Document row created — analysis targets CloudItem directly.
|
|
- Cache pins released on success, failure, and cancellation (T-14-12).
|
|
- Final version/fingerprint re-check marks items stale if metadata changed
|
|
since the job was enqueued (T-14-13).
|
|
|
|
Status transitions:
|
|
queued → downloading → extracting → classifying → indexed
|
|
→ cancelled (cooperative cancellation)
|
|
→ failed (transient or terminal error)
|
|
→ stale (item metadata changed before processing completed)
|
|
|
|
Rules:
|
|
- Service raises ValueError or domain exceptions only — never HTTPException
|
|
(CLAUDE.md service-layer rule).
|
|
- content_hash is computed ONLY while bytes are already in memory — never
|
|
as a precondition for fetching bytes (D-20).
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import uuid
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
from typing import Optional, TYPE_CHECKING
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from db.models import (
|
|
CloudAnalysisJob,
|
|
CloudAnalysisJobItem,
|
|
CloudByteCacheEntry,
|
|
CloudItem,
|
|
)
|
|
from services.cloud_analysis_versioning import compute_version_key
|
|
|
|
if TYPE_CHECKING:
|
|
pass
|
|
|
|
|
|
# ── Domain exceptions ─────────────────────────────────────────────────────────
|
|
|
|
class ItemCancelled(ValueError):
|
|
"""Processing cancelled cooperatively — caller should stop without retrying."""
|
|
|
|
|
|
class ItemStale(ValueError):
|
|
"""Item metadata changed between enqueue and processing — mark stale."""
|
|
|
|
|
|
class OwnerValidationError(ValueError):
|
|
"""Connection or item no longer belongs to the expected user."""
|
|
|
|
|
|
# ── Processing result ─────────────────────────────────────────────────────────
|
|
|
|
@dataclass
|
|
class ProcessingResult:
|
|
"""Return value from process_job_item.
|
|
|
|
Attributes:
|
|
status: Final CloudAnalysisJobItem status ("indexed", "cancelled",
|
|
"failed", "stale").
|
|
cache_entry_id: UUID of the cache entry used, or None.
|
|
error_code: Short machine-readable error code on failure, or None.
|
|
error_message: Human-readable detail on failure, or None.
|
|
"""
|
|
status: str
|
|
cache_entry_id: Optional[uuid.UUID] = None
|
|
error_code: Optional[str] = None
|
|
error_message: Optional[str] = None
|
|
|
|
|
|
# ── Transition helper ─────────────────────────────────────────────────────────
|
|
|
|
async def _set_item_status(
|
|
session: AsyncSession,
|
|
*,
|
|
job_item: CloudAnalysisJobItem,
|
|
job: CloudAnalysisJob,
|
|
new_status: str,
|
|
old_status: str,
|
|
error_code: Optional[str] = None,
|
|
error_message: Optional[str] = None,
|
|
cache_entry_id: Optional[uuid.UUID] = None,
|
|
) -> None:
|
|
"""Update job item status and adjust aggregate job counters atomically.
|
|
|
|
Decrement the old-status counter and increment the new-status counter on
|
|
the job row. Writes finished_at when transitioning to a terminal state.
|
|
"""
|
|
now = datetime.now(timezone.utc)
|
|
terminal_statuses = {"indexed", "failed", "cancelled", "stale", "unsupported", "already_current"}
|
|
|
|
job_item.status = new_status
|
|
job_item.updated_at = now
|
|
if error_code is not None:
|
|
job_item.error_code = error_code
|
|
if error_message is not None:
|
|
job_item.error_message = error_message
|
|
if cache_entry_id is not None:
|
|
job_item.cache_entry_id = cache_entry_id
|
|
if new_status in terminal_statuses:
|
|
job_item.finished_at = now
|
|
|
|
# Counter adjustment map: status name → attribute name on CloudAnalysisJob
|
|
_STATUS_COUNTER = {
|
|
"queued": "queued_count",
|
|
"downloading": "downloading_count",
|
|
"extracting": "extracting_count",
|
|
"classifying": "classifying_count",
|
|
"indexed": "indexed_count",
|
|
"already_current": "already_current_count",
|
|
"cancelled": "cancelled_count",
|
|
"failed": "failed_count",
|
|
"unsupported": "unsupported_count",
|
|
"stale": "failed_count", # stale counts as a failed outcome for UI counters
|
|
}
|
|
|
|
old_attr = _STATUS_COUNTER.get(old_status)
|
|
new_attr = _STATUS_COUNTER.get(new_status)
|
|
|
|
if old_attr:
|
|
current_old = getattr(job, old_attr, 0) or 0
|
|
setattr(job, old_attr, max(0, current_old - 1))
|
|
if new_attr and new_attr != old_attr:
|
|
current_new = getattr(job, new_attr, 0) or 0
|
|
setattr(job, new_attr, current_new + 1)
|
|
|
|
job.updated_at = now
|
|
|
|
# Transition job to running if not already in an active or terminal state
|
|
active_or_terminal = {"running", "completed", "cancelled", "failed"}
|
|
if job.status not in active_or_terminal:
|
|
job.status = "running"
|
|
if job.started_at is None:
|
|
job.started_at = now
|
|
|
|
await session.flush()
|
|
|
|
|
|
async def _update_job_completion(
|
|
session: AsyncSession,
|
|
*,
|
|
job: CloudAnalysisJob,
|
|
) -> None:
|
|
"""Check if all job items are terminal and advance job to completed/failed/cancelled."""
|
|
now = datetime.now(timezone.utc)
|
|
|
|
# Items still pending (in non-terminal states)
|
|
active_count = (
|
|
(job.queued_count or 0)
|
|
+ (job.downloading_count or 0)
|
|
+ (job.extracting_count or 0)
|
|
+ (job.classifying_count or 0)
|
|
)
|
|
|
|
if active_count == 0 and job.status not in ("completed", "cancelled", "failed"):
|
|
# All items terminal — determine final job outcome
|
|
if (job.failed_count or 0) > 0:
|
|
job.status = "failed"
|
|
elif (job.cancelled_count or 0) > 0 and (job.indexed_count or 0) == 0:
|
|
job.status = "cancelled"
|
|
else:
|
|
job.status = "completed"
|
|
job.finished_at = now
|
|
job.updated_at = now
|
|
await session.flush()
|
|
|
|
|
|
# ── Main processing entry point ───────────────────────────────────────────────
|
|
|
|
async def process_job_item(
|
|
session: AsyncSession,
|
|
*,
|
|
job_id: uuid.UUID,
|
|
item_id: uuid.UUID,
|
|
user_id: uuid.UUID,
|
|
connection_id: uuid.UUID,
|
|
cloud_item_id: uuid.UUID,
|
|
minio_client,
|
|
provider_adapter,
|
|
) -> ProcessingResult:
|
|
"""Process a single CloudAnalysisJobItem: hydrate bytes, extract, classify.
|
|
|
|
This function:
|
|
- Never calls provider mutation methods (upload/delete/rename/move/create_folder).
|
|
- Never creates a local Document row.
|
|
- Always releases cache pins in a finally block.
|
|
- Returns a ProcessingResult describing the final status.
|
|
|
|
Args:
|
|
session: Active async SQLAlchemy session (per-call, never shared).
|
|
job_id: CloudAnalysisJob UUID.
|
|
item_id: CloudAnalysisJobItem UUID.
|
|
user_id: Owner UUID — must match job + item + connection.
|
|
connection_id: CloudConnection UUID.
|
|
cloud_item_id: CloudItem UUID.
|
|
minio_client: MinIO client for cache byte storage (get_object/put_object).
|
|
provider_adapter: CloudResourceAdapter instance for reading provider bytes.
|
|
Must NOT expose mutation methods to this function.
|
|
|
|
Returns:
|
|
ProcessingResult with final status and metadata.
|
|
"""
|
|
uid = user_id if isinstance(user_id, uuid.UUID) else uuid.UUID(str(user_id))
|
|
jid = job_id if isinstance(job_id, uuid.UUID) else uuid.UUID(str(job_id))
|
|
iid = item_id if isinstance(item_id, uuid.UUID) else uuid.UUID(str(item_id))
|
|
cid = connection_id if isinstance(connection_id, uuid.UUID) else uuid.UUID(str(connection_id))
|
|
ciid = cloud_item_id if isinstance(cloud_item_id, uuid.UUID) else uuid.UUID(str(cloud_item_id))
|
|
|
|
cache_entry_id: Optional[uuid.UUID] = None
|
|
cache_pinned = False
|
|
|
|
try:
|
|
# ── Step 1: Revalidate owner / connection / item ──────────────────────
|
|
from services.cloud_items import resolve_owned_connection, ConnectionNotFound
|
|
|
|
try:
|
|
conn = await resolve_owned_connection(session, connection_id=cid, user_id=uid)
|
|
except ConnectionNotFound:
|
|
return ProcessingResult(
|
|
status="failed",
|
|
error_code="connection_not_found",
|
|
error_message="Connection no longer exists or does not belong to this user.",
|
|
)
|
|
|
|
# Reload the job and job item (revalidate they still exist and belong to user)
|
|
job_result = await session.execute(
|
|
select(CloudAnalysisJob).where(
|
|
CloudAnalysisJob.id == jid,
|
|
CloudAnalysisJob.user_id == uid,
|
|
)
|
|
)
|
|
job = job_result.scalars().first()
|
|
if job is None:
|
|
return ProcessingResult(
|
|
status="failed",
|
|
error_code="job_not_found",
|
|
error_message="Analysis job not found or does not belong to this user.",
|
|
)
|
|
|
|
item_result = await session.execute(
|
|
select(CloudAnalysisJobItem).where(
|
|
CloudAnalysisJobItem.id == iid,
|
|
CloudAnalysisJobItem.job_id == jid,
|
|
CloudAnalysisJobItem.user_id == uid,
|
|
)
|
|
)
|
|
job_item = item_result.scalars().first()
|
|
if job_item is None:
|
|
return ProcessingResult(
|
|
status="failed",
|
|
error_code="item_not_found",
|
|
error_message="Job item not found.",
|
|
)
|
|
|
|
old_status = job_item.status
|
|
|
|
# Reload the CloudItem
|
|
cloud_item_result = await session.execute(
|
|
select(CloudItem).where(
|
|
CloudItem.id == ciid,
|
|
CloudItem.user_id == uid,
|
|
)
|
|
)
|
|
cloud_item = cloud_item_result.scalars().first()
|
|
if cloud_item is None:
|
|
# Item was deleted
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="cancelled",
|
|
old_status=old_status,
|
|
error_code="item_deleted",
|
|
error_message="Cloud item no longer exists.",
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
return ProcessingResult(
|
|
status="cancelled",
|
|
error_code="item_deleted",
|
|
error_message="Cloud item no longer exists.",
|
|
)
|
|
|
|
# ── Step 2: Cooperative cancellation check ────────────────────────────
|
|
if job.status == "cancelled" or job_item.status == "cancelled":
|
|
# Already cancelled — ensure status is correct and return
|
|
if job_item.status != "cancelled":
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="cancelled",
|
|
old_status=old_status,
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
return ProcessingResult(status="cancelled")
|
|
|
|
# ── Step 3: Recompute version key; detect stale condition ─────────────
|
|
modified_str: Optional[str] = None
|
|
if cloud_item.modified_at is not None:
|
|
modified_str = cloud_item.modified_at.isoformat()
|
|
|
|
current_vk = compute_version_key(
|
|
provider_item_id=cloud_item.provider_item_id,
|
|
version=getattr(cloud_item, "version", None),
|
|
etag=cloud_item.etag,
|
|
size=cloud_item.provider_size,
|
|
modified_at=modified_str,
|
|
content_type=cloud_item.content_type,
|
|
)
|
|
|
|
enqueued_vk = job_item.version_key
|
|
if enqueued_vk and current_vk != enqueued_vk:
|
|
# Item metadata changed between enqueue and processing — mark stale
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="stale",
|
|
old_status=old_status,
|
|
error_code="version_changed",
|
|
error_message=(
|
|
"Item metadata changed between enqueue and processing. "
|
|
"Re-enqueue to analyse the current version."
|
|
),
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
return ProcessingResult(
|
|
status="stale",
|
|
error_code="version_changed",
|
|
error_message="Item metadata changed between enqueue and processing.",
|
|
)
|
|
|
|
# ── Transition: queued → downloading ──────────────────────────────────
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="downloading",
|
|
old_status=old_status,
|
|
)
|
|
old_status = "downloading"
|
|
await session.commit()
|
|
|
|
# ── Step 4: Check cache; hydrate if needed ────────────────────────────
|
|
from services.cloud_cache import (
|
|
retain_or_reuse_cache_entry,
|
|
pin_cache_entry,
|
|
release_cache_entry,
|
|
increment_quota_for_cache,
|
|
CacheQuotaExceeded,
|
|
)
|
|
|
|
# Re-check cancellation after commit
|
|
await session.refresh(job)
|
|
await session.refresh(job_item)
|
|
if job.status == "cancelled" or job_item.status == "cancelled":
|
|
return ProcessingResult(status="cancelled")
|
|
|
|
file_bytes: Optional[bytes] = None
|
|
|
|
# Check for a valid (non-evicted) cache entry for this version key
|
|
from sqlalchemy import and_
|
|
cache_check = await session.execute(
|
|
select(CloudByteCacheEntry).where(
|
|
and_(
|
|
CloudByteCacheEntry.user_id == uid,
|
|
CloudByteCacheEntry.connection_id == cid,
|
|
CloudByteCacheEntry.cloud_item_id == ciid,
|
|
CloudByteCacheEntry.version_key == current_vk,
|
|
CloudByteCacheEntry.evicted_at.is_(None),
|
|
)
|
|
).limit(1)
|
|
)
|
|
existing_cache = cache_check.scalars().first()
|
|
|
|
if existing_cache is not None:
|
|
# Cache hit: pin and retrieve bytes from MinIO
|
|
await pin_cache_entry(session, entry_id=existing_cache.id, user_id=uid)
|
|
cache_pinned = True
|
|
cache_entry_id = existing_cache.id
|
|
await session.commit()
|
|
|
|
try:
|
|
file_bytes = await minio_client.get_object(existing_cache.object_key)
|
|
except Exception as exc:
|
|
# Cache miss recovery: fall through to provider download
|
|
await release_cache_entry(session, entry_id=existing_cache.id, user_id=uid)
|
|
cache_pinned = False
|
|
await session.commit()
|
|
existing_cache = None
|
|
file_bytes = None
|
|
|
|
if file_bytes is None:
|
|
# Cache miss: download from provider
|
|
try:
|
|
file_bytes = await _download_from_provider(
|
|
provider_adapter, cloud_item.provider_item_id
|
|
)
|
|
except Exception as exc:
|
|
err_str = str(exc).lower()
|
|
if any(kw in err_str for kw in ("unauthorized", "401", "403", "invalid_grant", "scope")):
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="failed",
|
|
old_status=old_status,
|
|
error_code="auth_error",
|
|
error_message="Provider authentication failed. Re-connect the account.",
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
await session.commit()
|
|
return ProcessingResult(
|
|
status="failed",
|
|
error_code="auth_error",
|
|
error_message="Provider authentication failed.",
|
|
)
|
|
# Transient provider error — re-raise so Celery can retry
|
|
raise
|
|
|
|
# Compute content hash while bytes are already in memory (D-20)
|
|
content_hash = hashlib.sha256(file_bytes).hexdigest()
|
|
size_bytes = len(file_bytes)
|
|
|
|
# Store bytes in MinIO cache
|
|
object_key = f"cache/{uid}/{uuid.uuid4()}{_ext_from_content_type(cloud_item.content_type)}"
|
|
try:
|
|
await minio_client.put_object(object_key, file_bytes, content_type=cloud_item.content_type)
|
|
except Exception as exc:
|
|
# Cache store failure is non-fatal — continue without caching
|
|
object_key = None
|
|
|
|
if object_key is not None:
|
|
# Update quota atomically
|
|
try:
|
|
await increment_quota_for_cache(session, user_id=uid, size_bytes=size_bytes)
|
|
except CacheQuotaExceeded:
|
|
# Quota exhausted — proceed without caching (analysis still possible)
|
|
# Try to delete the MinIO object we just stored
|
|
try:
|
|
await minio_client.delete_object(object_key)
|
|
except Exception:
|
|
pass
|
|
object_key = None
|
|
|
|
if object_key is not None:
|
|
# Create or reactivate cache entry
|
|
entry, _created = await retain_or_reuse_cache_entry(
|
|
session,
|
|
user_id=uid,
|
|
connection_id=cid,
|
|
cloud_item_id=ciid,
|
|
provider_item_id=cloud_item.provider_item_id,
|
|
version_key=current_vk,
|
|
object_key=object_key,
|
|
content_type=cloud_item.content_type,
|
|
size_bytes=size_bytes,
|
|
content_hash=content_hash,
|
|
)
|
|
await pin_cache_entry(session, entry_id=entry.id, user_id=uid)
|
|
cache_pinned = True
|
|
cache_entry_id = entry.id
|
|
await session.commit()
|
|
|
|
# ── Re-check cancellation after byte hydration ────────────────────────
|
|
await session.refresh(job)
|
|
await session.refresh(job_item)
|
|
if job.status == "cancelled" or job_item.status == "cancelled":
|
|
return ProcessingResult(status="cancelled")
|
|
|
|
# ── Transition: downloading → extracting ──────────────────────────────
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="extracting",
|
|
old_status=old_status,
|
|
)
|
|
old_status = "extracting"
|
|
await session.commit()
|
|
|
|
# ── Step 5: Extract text from bytes ───────────────────────────────────
|
|
from services.extractor import extract_text_from_bytes
|
|
|
|
try:
|
|
extracted_text = extract_text_from_bytes(
|
|
file_bytes, cloud_item.content_type or "application/octet-stream"
|
|
)
|
|
except Exception as exc:
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="failed",
|
|
old_status=old_status,
|
|
error_code="extraction_failed",
|
|
error_message=f"Text extraction failed: {type(exc).__name__}",
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
await session.commit()
|
|
return ProcessingResult(
|
|
status="failed",
|
|
cache_entry_id=cache_entry_id,
|
|
error_code="extraction_failed",
|
|
error_message=f"Text extraction failed: {type(exc).__name__}",
|
|
)
|
|
|
|
# ── Step 6: Persist extracted text onto CloudItem ─────────────────────
|
|
cloud_item.extracted_text = extracted_text
|
|
cloud_item.updated_at = datetime.now(timezone.utc)
|
|
await session.flush()
|
|
|
|
# Re-check cancellation after extraction
|
|
await session.refresh(job)
|
|
await session.refresh(job_item)
|
|
if job.status == "cancelled" or job_item.status == "cancelled":
|
|
return ProcessingResult(status="cancelled")
|
|
|
|
# ── Transition: extracting → classifying ──────────────────────────────
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="classifying",
|
|
old_status=old_status,
|
|
)
|
|
old_status = "classifying"
|
|
await session.commit()
|
|
|
|
# ── Step 7: Classify topics ───────────────────────────────────────────
|
|
try:
|
|
topics = await _classify_cloud_item(session, cloud_item=cloud_item)
|
|
except Exception as exc:
|
|
# Classification failure is retryable in the Celery layer
|
|
raise _ClassificationError(f"Classification failed: {exc}") from exc
|
|
|
|
# ── Step 8: Update CloudItem.analysis_status ──────────────────────────
|
|
now = datetime.now(timezone.utc)
|
|
cloud_item.analysis_status = "indexed"
|
|
cloud_item.updated_at = now
|
|
await session.flush()
|
|
|
|
# ── Final stale check: re-confirm version key still matches ───────────
|
|
# Re-read the item to get any updates from concurrent reconcile jobs
|
|
final_vk = compute_version_key(
|
|
provider_item_id=cloud_item.provider_item_id,
|
|
version=getattr(cloud_item, "version", None),
|
|
etag=cloud_item.etag,
|
|
size=cloud_item.provider_size,
|
|
modified_at=(cloud_item.modified_at.isoformat() if cloud_item.modified_at else None),
|
|
content_type=cloud_item.content_type,
|
|
)
|
|
if final_vk != current_vk:
|
|
# Provider metadata changed during processing — mark stale
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="stale",
|
|
old_status=old_status,
|
|
error_code="version_changed_during_processing",
|
|
error_message="Item metadata changed during processing. Re-enqueue to refresh.",
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
await session.commit()
|
|
return ProcessingResult(
|
|
status="stale",
|
|
cache_entry_id=cache_entry_id,
|
|
error_code="version_changed_during_processing",
|
|
)
|
|
|
|
# ── Transition: classifying → indexed ─────────────────────────────────
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="indexed",
|
|
old_status=old_status,
|
|
cache_entry_id=cache_entry_id,
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
await session.commit()
|
|
|
|
return ProcessingResult(
|
|
status="indexed",
|
|
cache_entry_id=cache_entry_id,
|
|
)
|
|
|
|
except _ClassificationError:
|
|
# Re-raise so the Celery task layer can call self.retry()
|
|
raise
|
|
|
|
except Exception:
|
|
# Unexpected error — mark failed without leaking exception text
|
|
try:
|
|
if "job_item" in dir() and "job" in dir() and "old_status" in dir():
|
|
await _set_item_status(
|
|
session,
|
|
job_item=job_item,
|
|
job=job,
|
|
new_status="failed",
|
|
old_status=old_status,
|
|
error_code="unexpected_error",
|
|
error_message="An unexpected error occurred during processing.",
|
|
)
|
|
await _update_job_completion(session, job=job)
|
|
await session.commit()
|
|
except Exception:
|
|
pass
|
|
raise
|
|
|
|
finally:
|
|
# Always release cache pin regardless of outcome (T-14-12)
|
|
if cache_pinned and cache_entry_id is not None:
|
|
try:
|
|
from services.cloud_cache import release_cache_entry
|
|
await release_cache_entry(session, entry_id=cache_entry_id, user_id=uid)
|
|
await session.commit()
|
|
except Exception:
|
|
pass # Pin release failure must not mask the original result
|
|
|
|
|
|
# ── Classification helper ─────────────────────────────────────────────────────
|
|
|
|
class _ClassificationError(Exception):
|
|
"""Sentinel for retryable classification failures — escapes asyncio.run()."""
|
|
|
|
|
|
async def _classify_cloud_item(
|
|
session: AsyncSession,
|
|
*,
|
|
cloud_item: CloudItem,
|
|
) -> list[str]:
|
|
"""Classify a CloudItem using the AI provider pipeline.
|
|
|
|
Adapted from services.classifier.classify_document but targets CloudItem
|
|
directly — no Document row created or required.
|
|
|
|
Returns the list of assigned topic names.
|
|
"""
|
|
from db.models import User, Topic, CloudItemTopic
|
|
from sqlalchemy import select as sa_select
|
|
from services import storage as doc_storage
|
|
from services.ai_config import load_provider_config
|
|
from ai import get_provider
|
|
from ai.provider_config import ProviderConfig, PROVIDER_DEFAULTS
|
|
from config import settings as app_settings
|
|
|
|
_DEFAULT_SYSTEM_PROMPT = (
|
|
"You are a document classification assistant. When given a document's text "
|
|
"content and a list of existing topics, you must:\n"
|
|
"1. Assign the document to one or more relevant topics from the list.\n"
|
|
"2. If no existing topics fit well, suggest new topic names.\n"
|
|
"Return ONLY valid JSON in this exact format, with no additional text or "
|
|
'explanation:\n{"assigned_topics": ["topic1"], "new_topic_suggestions": '
|
|
'["new topic name"]}\n'
|
|
"If the document fits no topics and you have no suggestions, return: "
|
|
'{"assigned_topics": [], "new_topic_suggestions": []}'
|
|
)
|
|
|
|
if not cloud_item.extracted_text:
|
|
# Nothing to classify — skip silently
|
|
return []
|
|
|
|
# Load user AI preferences
|
|
user_result = await session.execute(
|
|
sa_select(User).where(User.id == cloud_item.user_id)
|
|
)
|
|
user = user_result.scalars().first()
|
|
|
|
ai_provider = (user.ai_provider if user else None)
|
|
ai_model = (user.ai_model if user else None)
|
|
|
|
# Resolve provider config (same as classifier.classify_document)
|
|
if ai_provider is not None:
|
|
config = ProviderConfig(
|
|
provider_id=ai_provider,
|
|
model=ai_model or PROVIDER_DEFAULTS.get(ai_provider, {}).get("model", ""),
|
|
api_key="",
|
|
base_url=None,
|
|
context_chars=PROVIDER_DEFAULTS.get(ai_provider, {}).get("context_chars", 8000),
|
|
)
|
|
else:
|
|
config = await load_provider_config(session)
|
|
if config is None:
|
|
fallback_provider = app_settings.default_ai_provider
|
|
config = ProviderConfig(
|
|
provider_id=fallback_provider,
|
|
model=app_settings.default_ai_model,
|
|
api_key="",
|
|
base_url=None,
|
|
context_chars=PROVIDER_DEFAULTS.get(fallback_provider, {}).get(
|
|
"context_chars", 8000
|
|
),
|
|
)
|
|
|
|
provider = get_provider(config)
|
|
system_prompt = app_settings.system_prompt or _DEFAULT_SYSTEM_PROMPT
|
|
|
|
# Load user's topics for namespace-scoped classification (D-17)
|
|
topics_data = await doc_storage.load_topics_for_user(
|
|
session, user_id=cloud_item.user_id
|
|
)
|
|
topic_names = [t["name"] for t in topics_data]
|
|
|
|
result = await provider.classify(
|
|
cloud_item.extracted_text, topic_names, system_prompt
|
|
)
|
|
|
|
# Auto-create suggested topics in the user's namespace (D-11)
|
|
existing_names = {t.lower() for t in topic_names}
|
|
all_new_names = set(result.suggested_new_topics) | set(result.topics)
|
|
for name in all_new_names:
|
|
if name.strip() and name.lower() not in existing_names:
|
|
await doc_storage.create_topic(
|
|
session, name.strip(), user_id=cloud_item.user_id
|
|
)
|
|
existing_names.add(name.lower())
|
|
|
|
# Build the final topic list
|
|
final_topics = [
|
|
t for t in list(set(result.topics + result.suggested_new_topics))
|
|
if t.strip()
|
|
]
|
|
|
|
# Associate topics with the CloudItem (not Document)
|
|
if final_topics:
|
|
# Fetch or create Topic rows
|
|
topic_rows_result = await session.execute(
|
|
sa_select(Topic).where(
|
|
Topic.name.in_(final_topics),
|
|
Topic.user_id == cloud_item.user_id,
|
|
)
|
|
)
|
|
topic_rows = topic_rows_result.scalars().all()
|
|
existing_topic_map = {t.name.lower(): t for t in topic_rows}
|
|
|
|
now = datetime.now(timezone.utc)
|
|
for topic_name in final_topics:
|
|
topic_row = existing_topic_map.get(topic_name.lower())
|
|
if topic_row is None:
|
|
continue # create_topic should have created it above
|
|
|
|
# Upsert CloudItemTopic association
|
|
existing_assoc = await session.execute(
|
|
sa_select(CloudItemTopic).where(
|
|
CloudItemTopic.cloud_item_id == cloud_item.id,
|
|
CloudItemTopic.topic_id == topic_row.id,
|
|
)
|
|
)
|
|
if existing_assoc.scalars().first() is None:
|
|
assoc = CloudItemTopic(
|
|
cloud_item_id=cloud_item.id,
|
|
topic_id=topic_row.id,
|
|
)
|
|
session.add(assoc)
|
|
|
|
await session.flush()
|
|
|
|
return final_topics
|
|
|
|
|
|
# ── Provider download helper ──────────────────────────────────────────────────
|
|
|
|
async def _download_from_provider(adapter, provider_item_id: str) -> bytes:
|
|
"""Retrieve raw bytes for a cloud item.
|
|
|
|
Uses the content-read adapter path only. Mutation methods are never called.
|
|
The adapter interface is whatever the caller passes — test fakes and real
|
|
adapters are both accepted.
|
|
|
|
Args:
|
|
adapter: Provider adapter with a get_object(item_id) coroutine.
|
|
provider_item_id: Provider-native item identifier.
|
|
|
|
Returns:
|
|
Raw file bytes.
|
|
|
|
Raises:
|
|
Exception: Propagated from the adapter on failure.
|
|
"""
|
|
return await adapter.get_object(provider_item_id)
|
|
|
|
|
|
# ── Extension helper ──────────────────────────────────────────────────────────
|
|
|
|
def _ext_from_content_type(content_type: Optional[str]) -> str:
|
|
"""Return a file extension for caching MinIO objects."""
|
|
mapping = {
|
|
"application/pdf": ".pdf",
|
|
"application/msword": ".doc",
|
|
"application/vnd.openxmlformats-officedocument.wordprocessingml.document": ".docx",
|
|
"text/plain": ".txt",
|
|
"text/markdown": ".md",
|
|
"image/png": ".png",
|
|
"image/jpeg": ".jpg",
|
|
"image/gif": ".gif",
|
|
"image/webp": ".webp",
|
|
"image/tiff": ".tiff",
|
|
}
|
|
return mapping.get(content_type or "", ".bin")
|