Files
dubizzle/iaai_scraper/storage/db.py
2026-04-09 20:35:25 +03:00

147 lines
5.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import logging
from contextlib import contextmanager
from datetime import datetime, timezone
from typing import Iterator
from sqlalchemy import create_engine, or_, select
from sqlalchemy.orm import Session, sessionmaker
from ..core.config import Settings
from .models import Base, Car, Image, SyncRun
from .schemas import CarRecord
logger = logging.getLogger("iaai_scraper.db")
CAR_DB_FIELDS = {
col.key for col in Car.__table__.columns
if col.key not in ("id",)
}
class PersistenceService:
def __init__(self, settings: Settings) -> None:
self.settings = settings
engine_kwargs = {
"echo": settings.database.echo,
"future": True,
}
if "postgresql" in settings.database.url:
engine_kwargs["pool_size"] = settings.database.pool_size
engine_kwargs["max_overflow"] = settings.database.max_overflow
engine_kwargs["pool_pre_ping"] = True
engine_kwargs["pool_recycle"] = settings.database.pool_recycle_seconds
self.engine = create_engine(settings.database.url, **engine_kwargs)
self.session_factory = sessionmaker(bind=self.engine, expire_on_commit=False, future=True)
def create_tables(self) -> None:
# В тестах/локально на SQLite разрешаем create_all; для non-SQLite в проде — только через миграции.
is_sqlite = self.settings.database.url.startswith("sqlite")
if not is_sqlite and not self.settings.database.auto_create_tables:
return
try:
Base.metadata.create_all(self.engine)
except Exception:
logger.debug("create_tables skipped (schema already exists)")
@contextmanager
def session_scope(self) -> Iterator[Session]:
session = self.session_factory()
try:
yield session
session.commit()
except Exception:
session.rollback()
raise
finally:
session.close()
def start_sync_run(self, lane: str) -> int:
with self.session_scope() as session:
now = datetime.now(timezone.utc)
stale_runs = session.execute(select(SyncRun).where(SyncRun.status == "running")).scalars().all()
for stale in stale_runs:
stale.status = "failed"
stale.finished_at = now
if not stale.error_summary:
stale.error_summary = "Recovered stale running sync run before starting a new run"
run = SyncRun(status="running", lane=lane, ids_fetched=0, cars_upserted=0, cars_failed=0, images_upserted=0)
session.add(run)
session.flush()
return int(run.id)
def finish_sync_run(self, run_id: int, *, status: str, ids_fetched: int, cars_upserted: int, cars_failed: int, images_upserted: int, error_summary: str | None = None) -> None:
with self.session_scope() as session:
run = session.get(SyncRun, run_id)
if run is None:
return
run.finished_at = datetime.now(timezone.utc)
run.status = status
run.ids_fetched = ids_fetched
run.cars_upserted = cars_upserted
run.cars_failed = cars_failed
run.images_upserted = images_upserted
run.error_summary = error_summary
def get_existing_origin_urls(self, origin_urls: list[str]) -> set[str]:
if not origin_urls:
return set()
with self.session_scope() as session:
rows = session.execute(select(Car.origin_url).where(Car.origin_url.in_(origin_urls))).all()
return {str(row[0]) for row in rows if row and row[0]}
def get_existing_origin_ids(self, origin_ids: list[str]) -> set[str]:
if not origin_ids:
return set()
with self.session_scope() as session:
rows = session.execute(select(Car.origin_id).where(Car.origin_id.in_(origin_ids))).all()
return {str(row[0]) for row in rows if row and row[0]}
@staticmethod
def _add_images(session: Session, car_id: int, images: list[dict[str, object]]) -> None:
for image_payload in images:
session.add(Image(fullres_image=str(image_payload["fullres_image"]), preview_image=str(image_payload["preview_image"]), order_index=int(image_payload.get("order_index", 0)), car_id=car_id))
@staticmethod
def _car_payload(record: CarRecord) -> dict[str, object]:
payload = record.model_dump(mode="python")
return {key: value for key, value in payload.items() if key in CAR_DB_FIELDS}
def upsert_car(self, record: CarRecord):
# Insert/update автомобиля по origin_id/origin_url.
payload = self._car_payload(record)
images = [image.model_dump(mode="python") for image in record.images]
with self.session_scope() as session:
car = session.execute(
select(Car).where(or_(Car.origin_id == record.origin_id, Car.origin_url == record.origin_url))
).scalar_one_or_none()
action = "inserted"
if car is None:
car = Car(**payload)
session.add(car)
session.flush()
else:
action = "updated"
for key, value in payload.items():
setattr(car, key, value)
car.last_seen_at = record.last_seen_at
session.flush()
nested = session.begin_nested()
try:
for image in list(car.images):
session.delete(image)
session.flush()
self._add_images(session, int(car.id), images)
session.flush()
nested.commit()
except Exception:
nested.rollback()
logger.warning("Image replacement failed for car %s, keeping old images", record.origin_id)
images = []
return {"car_id": int(car.id), "images_upserted": len(images), "action": action}
self._add_images(session, int(car.id), images)
session.flush()
return {"car_id": int(car.id), "images_upserted": len(images), "action": action}