feat(14-05): add cloud analysis processing service (Task 1)
- 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
This commit is contained in:
@@ -0,0 +1,822 @@
|
||||
"""
|
||||
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")
|
||||
Reference in New Issue
Block a user