314 lines
11 KiB
Python
314 lines
11 KiB
Python
import asyncio
|
|
import uuid
|
|
from datetime import UTC, datetime, timedelta
|
|
|
|
import httpx
|
|
import redis.asyncio as redis
|
|
import structlog
|
|
from sqlalchemy import delete, select
|
|
|
|
from app.db import Database, DeliveryOutbox, Dialog, Message, MessageAttachment, SafetyTask
|
|
from app.integrations import (
|
|
DependencyFailure,
|
|
OpenLinesClient,
|
|
S3Client,
|
|
SafetyClient,
|
|
fresh_openlines_payload,
|
|
)
|
|
from app.notification_models import ClientUploadDraft
|
|
from app.notification_service import expire_notifications
|
|
from app.realtime import RealtimeFanout
|
|
from app.services import load_settings, publish_dialog_status, publish_message_status
|
|
from app.settings import Settings, get_settings
|
|
|
|
log = structlog.get_logger()
|
|
|
|
|
|
async def delivery_once(
|
|
db: Database,
|
|
client: OpenLinesClient,
|
|
s3: S3Client,
|
|
fanout: RealtimeFanout,
|
|
settings: Settings,
|
|
worker_id: str,
|
|
batch_size: int = 20,
|
|
) -> int:
|
|
async with db.sessions() as session:
|
|
rows = (
|
|
(
|
|
await session.execute(
|
|
select(DeliveryOutbox)
|
|
.where(
|
|
DeliveryOutbox.status.in_(["pending", "retry"]),
|
|
DeliveryOutbox.next_attempt_at <= datetime.now(UTC),
|
|
)
|
|
.with_for_update(skip_locked=True)
|
|
.limit(batch_size)
|
|
)
|
|
)
|
|
.scalars()
|
|
.all()
|
|
)
|
|
ids = [row.id for row in rows]
|
|
for row in rows:
|
|
row.status = "processing"
|
|
row.locked_at = datetime.now(UTC)
|
|
row.locked_by = worker_id
|
|
await session.commit()
|
|
for row_id in ids:
|
|
async with db.sessions() as session:
|
|
row = await session.get(DeliveryOutbox, row_id, with_for_update=True)
|
|
if row is None:
|
|
continue
|
|
message = None
|
|
dialog = None
|
|
try:
|
|
payload = await fresh_openlines_payload(row.payload_json, s3)
|
|
await client.send(row.message_id, payload, f"worker-{worker_id}")
|
|
row.status = "delivered"
|
|
message = await session.get(Message, row.message_id)
|
|
if message:
|
|
message.delivery_status = "delivered"
|
|
dialog = await session.get(Dialog, message.dialog_id)
|
|
if dialog:
|
|
dialog.status = "waiting_for_company"
|
|
dialog.last_message_at = datetime.now(UTC)
|
|
except DependencyFailure:
|
|
row.attempt_count += 1
|
|
row.status = "dead_letter" if row.attempt_count >= 12 else "retry"
|
|
row.next_attempt_at = datetime.now(UTC) + timedelta(
|
|
seconds=min(3600, 2**row.attempt_count)
|
|
)
|
|
row.last_error_code = "dependency_unavailable"
|
|
row.locked_at = None
|
|
row.locked_by = None
|
|
await session.commit()
|
|
if message:
|
|
await publish_message_status(fanout, message, settings)
|
|
if dialog:
|
|
await publish_dialog_status(fanout, dialog)
|
|
return len(ids)
|
|
|
|
|
|
async def safety_once(
|
|
db: Database,
|
|
safety: SafetyClient,
|
|
s3: S3Client,
|
|
fanout: RealtimeFanout,
|
|
settings: Settings,
|
|
worker_id: str,
|
|
batch_size: int = 20,
|
|
) -> int:
|
|
async with db.sessions() as session:
|
|
rows = (
|
|
(
|
|
await session.execute(
|
|
select(SafetyTask)
|
|
.where(
|
|
SafetyTask.status.in_(["polling", "failed"]),
|
|
SafetyTask.next_poll_at <= datetime.now(UTC),
|
|
SafetyTask.deadline_at > datetime.now(UTC),
|
|
)
|
|
.with_for_update(skip_locked=True)
|
|
.limit(batch_size)
|
|
)
|
|
)
|
|
.scalars()
|
|
.all()
|
|
)
|
|
ids = [row.id for row in rows]
|
|
for row in rows:
|
|
row.locked_at, row.locked_by = datetime.now(UTC), worker_id
|
|
await session.commit()
|
|
for task_id in ids:
|
|
async with db.sessions() as session:
|
|
task = await session.get(SafetyTask, task_id, with_for_update=True)
|
|
if task is None:
|
|
continue
|
|
message = None
|
|
try:
|
|
verdict = await safety.poll(task.task_id, f"worker-{worker_id}")
|
|
message = await session.get(Message, task.message_id)
|
|
attachment = (
|
|
await session.get(MessageAttachment, task.attachment_id)
|
|
if task.attachment_id
|
|
else None
|
|
)
|
|
if verdict["_status"] == 200 and message:
|
|
if attachment and attachment.quarantine_object_key:
|
|
destination = f"attachments/dialogs/{message.dialog_id}/{attachment.id}"
|
|
await s3.promote(attachment.quarantine_object_key, destination)
|
|
attachment.storage_bucket = s3.settings.selectel_s3_bucket_attachments
|
|
attachment.object_key = destination
|
|
attachment.quarantine_object_key = None
|
|
attachment.scan_status = "clean"
|
|
message.safety_status = "allowed"
|
|
task.status = "completed"
|
|
elif verdict["_status"] == 403 and message:
|
|
message.text = ""
|
|
message.safety_status = "blocked"
|
|
message.delivery_status = "rejected"
|
|
if attachment and attachment.quarantine_object_key:
|
|
await s3.delete_quarantine(attachment.quarantine_object_key)
|
|
attachment.scan_status = "infected"
|
|
task.status = "completed"
|
|
else:
|
|
task.next_poll_at = datetime.now(UTC) + timedelta(seconds=2)
|
|
except DependencyFailure:
|
|
task.attempt_count += 1
|
|
task.status = "failed"
|
|
task.next_poll_at = datetime.now(UTC) + timedelta(
|
|
seconds=min(300, 2**task.attempt_count)
|
|
)
|
|
task.locked_at = None
|
|
task.locked_by = None
|
|
await session.commit()
|
|
if message:
|
|
await publish_message_status(fanout, message, settings)
|
|
return len(ids)
|
|
|
|
|
|
async def cleanup_once(db: Database, s3: S3Client, batch_size: int = 100) -> int:
|
|
async with db.sessions() as session:
|
|
rows = (
|
|
(
|
|
await session.execute(
|
|
select(MessageAttachment)
|
|
.where(
|
|
MessageAttachment.record_status == "A",
|
|
MessageAttachment.quarantine_object_key.is_not(None),
|
|
MessageAttachment.upload_expires_at < datetime.now(UTC),
|
|
MessageAttachment.message_id.is_(None),
|
|
)
|
|
.with_for_update(skip_locked=True)
|
|
.limit(batch_size)
|
|
)
|
|
)
|
|
.scalars()
|
|
.all()
|
|
)
|
|
for row in rows:
|
|
if row.quarantine_object_key:
|
|
await s3.delete_quarantine(row.quarantine_object_key)
|
|
row.record_status = "D"
|
|
row.status_changed_at = datetime.now(UTC)
|
|
row.status_change_reason = "expired_quarantine_cleanup"
|
|
await session.commit()
|
|
return len(rows)
|
|
|
|
|
|
async def loop(kind: str) -> None:
|
|
settings = get_settings()
|
|
db = Database(settings.database_url)
|
|
http = httpx.AsyncClient()
|
|
safety = SafetyClient(settings, http)
|
|
openlines = OpenLinesClient(settings, http)
|
|
s3 = S3Client(settings)
|
|
redis_rt = redis.from_url(settings.redis_realtime_url, decode_responses=True)
|
|
fanout = RealtimeFanout(redis_rt)
|
|
worker_id = f"{kind}-{uuid.uuid4()}"
|
|
try:
|
|
while True:
|
|
count = 0
|
|
if kind == "delivery":
|
|
count = await delivery_once(db, openlines, s3, fanout, settings, worker_id)
|
|
elif kind == "safety":
|
|
count = await safety_once(db, safety, s3, fanout, settings, worker_id)
|
|
else:
|
|
count = await cleanup_once(db, s3)
|
|
if not count:
|
|
await asyncio.sleep(settings.worker_poll_interval_sec)
|
|
finally:
|
|
await http.aclose()
|
|
await redis_rt.aclose()
|
|
await db.close()
|
|
|
|
|
|
async def notification_expire_loop() -> None:
|
|
settings = get_settings()
|
|
db = Database(settings.database_url)
|
|
try:
|
|
while True:
|
|
async with db.sessions() as session:
|
|
snapshot = await load_settings(session)
|
|
run_at = snapshot.values["notification.expire_job.run_at"]
|
|
hour, minute = (int(value) for value in run_at.split(":", 1))
|
|
now = datetime.now(UTC)
|
|
target = now.replace(hour=hour, minute=minute, second=0, microsecond=0)
|
|
if target <= now:
|
|
target += timedelta(days=1)
|
|
await asyncio.sleep((target - now).total_seconds())
|
|
async with db.sessions() as session:
|
|
personal, guest = await expire_notifications(session)
|
|
log.info(
|
|
"notification.expired_batch",
|
|
personal_count=personal,
|
|
guest_count=guest,
|
|
)
|
|
finally:
|
|
await db.close()
|
|
|
|
|
|
async def notification_draft_cleanup_once(db: Database, s3: S3Client) -> int:
|
|
async with db.sessions() as session:
|
|
snapshot = await load_settings(session)
|
|
cutoff = datetime.now(UTC) - timedelta(
|
|
days=snapshot.integer("notification.upload_draft.ttl_days")
|
|
)
|
|
rows = (
|
|
(
|
|
await session.execute(
|
|
select(ClientUploadDraft)
|
|
.where(ClientUploadDraft.created_at < cutoff)
|
|
.with_for_update(skip_locked=True)
|
|
.limit(100)
|
|
)
|
|
)
|
|
.scalars()
|
|
.all()
|
|
)
|
|
for row in rows:
|
|
if row.state != "submitted":
|
|
if row.quarantine_object_key:
|
|
await s3.delete_quarantine(row.quarantine_object_key)
|
|
elif row.object_key:
|
|
await s3.delete(row.storage_bucket, row.object_key)
|
|
await session.execute(
|
|
delete(ClientUploadDraft).where(ClientUploadDraft.id == row.id)
|
|
)
|
|
await session.commit()
|
|
return len(rows)
|
|
|
|
|
|
async def notification_draft_cleanup_loop() -> None:
|
|
settings = get_settings()
|
|
db = Database(settings.database_url)
|
|
s3 = S3Client(settings)
|
|
try:
|
|
while True:
|
|
count = await notification_draft_cleanup_once(db, s3)
|
|
if count < 100:
|
|
await asyncio.sleep(86400)
|
|
finally:
|
|
await db.close()
|
|
|
|
|
|
def delivery_main() -> None:
|
|
asyncio.run(loop("delivery"))
|
|
|
|
|
|
def safety_main() -> None:
|
|
asyncio.run(loop("safety"))
|
|
|
|
|
|
def cleanup_main() -> None:
|
|
asyncio.run(loop("cleanup"))
|
|
|
|
|
|
def notification_expire_main() -> None:
|
|
asyncio.run(notification_expire_loop())
|
|
|
|
|
|
def notification_draft_cleanup_main() -> None:
|
|
asyncio.run(notification_draft_cleanup_loop())
|