from typing import Optional from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.models.company import Company, CompanyStatus from app.schemas.company import CompanyCreate, CompanyUpdate class CompanyService: """企业服务""" @staticmethod async def create_company( db: AsyncSession, company_data: CompanyCreate, ) -> Company: """ 创建企业 Args: db: 数据库会话 company_data: 企业数据 Returns: 创建的企业对象 """ company = Company(**company_data.model_dump()) db.add(company) await db.commit() await db.refresh(company) return company @staticmethod async def get_company( db: AsyncSession, company_id: int, ) -> Optional[Company]: """ 获取企业信息 Args: db: 数据库会话 company_id: 企业 ID Returns: 企业对象,不存在则返回 None """ result = await db.execute( select(Company).where( Company.id == company_id, Company.status != CompanyStatus.DELETED.value, ) ) return result.scalar_one_or_none() @staticmethod async def update_company( db: AsyncSession, company_id: int, company_data: CompanyUpdate, ) -> Optional[Company]: """ 更新企业信息 Args: db: 数据库会话 company_id: 企业 ID company_data: 更新数据 Returns: 更新后的企业对象,不存在则返回 None """ company = await CompanyService.get_company(db, company_id) if not company: return None update_data = company_data.model_dump(exclude_unset=True) for field, value in update_data.items(): setattr(company, field, value) await db.commit() await db.refresh(company) return company @staticmethod async def list_companies( db: AsyncSession, skip: int = 0, limit: int = 100, ) -> list[Company]: """ 获取企业列表 Args: db: 数据库会话 skip: 跳过数量 limit: 返回数量 Returns: 企业列表 """ result = await db.execute( select(Company) .where(Company.status != CompanyStatus.DELETED.value) .offset(skip) .limit(limit) ) return list(result.scalars().all()) # 单例实例 company_service = CompanyService()