90 lines
2.7 KiB
Python
90 lines
2.7 KiB
Python
from typing import List, Optional
|
|
|
|
from sqlalchemy import text
|
|
from sqlalchemy.orm import Session
|
|
|
|
from . import dto, models
|
|
|
|
DRIVER_ID_SEQ = "a76.driver_driver_id_seq"
|
|
|
|
|
|
def allocate_driver_id(db: Session) -> int:
|
|
"""Next surrogate driver_id (sequence from migration ca7d3c4e8b2a)."""
|
|
return db.execute(text(f"SELECT nextval('{DRIVER_ID_SEQ}')")).scalar()
|
|
|
|
|
|
class DriverService:
|
|
@staticmethod
|
|
def list_drivers(
|
|
db: Session, company_id: str, tenant_id: Optional[str] = None
|
|
) -> List[models.Driver]:
|
|
query = db.query(models.Driver).filter(models.Driver.company_id == company_id)
|
|
if tenant_id:
|
|
query = query.filter(models.Driver.tenant_id == tenant_id)
|
|
return query.all()
|
|
|
|
@staticmethod
|
|
def get_driver_by_key_and_line(
|
|
db: Session,
|
|
transporter_key: str,
|
|
line: int,
|
|
company_id: str,
|
|
tenant_id: Optional[str] = None,
|
|
) -> Optional[models.Driver]:
|
|
query = db.query(models.Driver).filter(
|
|
models.Driver.transporter_key == transporter_key,
|
|
models.Driver.line == line,
|
|
models.Driver.company_id == company_id,
|
|
)
|
|
if tenant_id:
|
|
query = query.filter(models.Driver.tenant_id == tenant_id)
|
|
return query.first()
|
|
|
|
@staticmethod
|
|
def create_driver(db: Session, driver_data: dto.DriverCreateDTO):
|
|
data = driver_data.model_dump()
|
|
if data.get("driver_id") is None:
|
|
data["driver_id"] = allocate_driver_id(db)
|
|
new_driver = models.Driver(**data)
|
|
db.add(new_driver)
|
|
db.commit()
|
|
db.refresh(new_driver)
|
|
return new_driver
|
|
|
|
@staticmethod
|
|
def update_driver(
|
|
db: Session,
|
|
transporter_key: str,
|
|
line: int,
|
|
company_id: str,
|
|
tenant_id: Optional[str],
|
|
data: dto.DriverUpdateDTO,
|
|
) -> Optional[models.Driver]:
|
|
driver = DriverService.get_driver_by_key_and_line(
|
|
db, transporter_key, line, company_id, tenant_id
|
|
)
|
|
if not driver:
|
|
return None
|
|
update_data = data.model_dump(exclude_unset=True)
|
|
for key, value in update_data.items():
|
|
setattr(driver, key, value)
|
|
db.commit()
|
|
db.refresh(driver)
|
|
return driver
|
|
|
|
@staticmethod
|
|
def delete_driver(
|
|
db: Session,
|
|
transporter_key: str,
|
|
line: int,
|
|
company_id: str,
|
|
tenant_id: Optional[str] = None,
|
|
) -> Optional[models.Driver]:
|
|
driver = DriverService.get_driver_by_key_and_line(
|
|
db, transporter_key, line, company_id, tenant_id
|
|
)
|
|
if driver:
|
|
db.delete(driver)
|
|
db.commit()
|
|
return driver
|