from typing import List, Optional from sqlalchemy import text from sqlalchemy.orm import Session from . import dto, models from api.v1.modules.a76.transportation.catalog_parity import ( driver_fields_to_csv_row, validate_driver_row_for_api, ) 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() validate_driver_row_for_api( int(driver_data.tenant_id), int(driver_data.company_id), driver_fields_to_csv_row(data), is_update=False, existing_driver_keys=set(), ) 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) merged = { "transporter_key": driver.transporter_key, "line": driver.line, "driver_name": driver.driver_name, "license_number": driver.license_number, "express_line_id": driver.express_line_id, "ace_id": driver.ace_id, "birth_date": driver.birth_date, "gender": driver.gender, "birth_country": driver.birth_country, "hazardous_material_auth": driver.hazardous_material_auth, "hazardous_material_state": driver.hazardous_material_state, "first_name": driver.first_name, "last_name": driver.last_name, "id_key1": driver.id_key1, "id_number1": driver.id_number1, "id_state1": driver.id_state1, "id_country1": driver.id_country1, "id_key2": driver.id_key2, "id_number2": driver.id_number2, "id_state2": driver.id_state2, "id_country2": driver.id_country2, "badge_number": driver.badge_number, "class_type": driver.class_type, "unique_badge_number": driver.unique_badge_number, } merged.update(update_data) tid = int(driver.tenant_id) cid = int(driver.company_id) validate_driver_row_for_api( tid, cid, driver_fields_to_csv_row(merged), is_update=True, existing_driver_keys={(driver.transporter_key.strip().upper(), driver.line)}, ) 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