feat(core): add address fields to company service serialization and CRUD
This commit is contained in:
@@ -16,6 +16,9 @@ from app.models.company import Company
|
||||
from app.models.contact import CompanyContact, Contact
|
||||
|
||||
|
||||
ADDRESS_FIELDS = ("address_street", "address_city", "address_zip", "address_country", "address_state")
|
||||
|
||||
|
||||
def _company_to_dict(c: Company, include_contacts: bool = False) -> dict[str, Any]:
|
||||
"""Serialize a Company ORM object to dict."""
|
||||
data: dict[str, Any] = {
|
||||
@@ -27,11 +30,16 @@ def _company_to_dict(c: Company, include_contacts: bool = False) -> dict[str, An
|
||||
"email": c.email,
|
||||
"website": c.website,
|
||||
"description": c.description,
|
||||
"address_street": c.address_street,
|
||||
"address_city": c.address_city,
|
||||
"address_zip": c.address_zip,
|
||||
"address_country": c.address_country,
|
||||
"address_state": c.address_state,
|
||||
"created_at": c.created_at.isoformat() if c.created_at else None,
|
||||
"updated_at": c.updated_at.isoformat() if c.updated_at else None,
|
||||
}
|
||||
if include_contacts:
|
||||
data["contacts"] = [] # populated by caller if needed
|
||||
data["contacts"] = []
|
||||
return data
|
||||
|
||||
|
||||
@@ -45,10 +53,7 @@ async def list_companies(
|
||||
sort_by: str = "name",
|
||||
sort_order: str = "asc",
|
||||
) -> dict[str, Any]:
|
||||
"""List companies with pagination, FTS search, industry filter, and sorting.
|
||||
|
||||
Returns {items, total, page, page_size}.
|
||||
"""
|
||||
"""List companies with pagination, FTS search, industry filter, and sorting."""
|
||||
page = max(1, page)
|
||||
page_size = max(1, min(100, page_size))
|
||||
|
||||
@@ -57,11 +62,9 @@ async def list_companies(
|
||||
Company.deleted_at.is_(None),
|
||||
)
|
||||
|
||||
# Industry filter
|
||||
if industry:
|
||||
base = base.where(Company.industry == industry)
|
||||
|
||||
# FTS search using PostgreSQL tsvector @@ plainto_tsquery, with ILIKE fallback
|
||||
if search:
|
||||
pattern = f"%{search}%"
|
||||
base = base.where(
|
||||
@@ -73,19 +76,16 @@ async def list_companies(
|
||||
)
|
||||
)
|
||||
|
||||
# Sorting
|
||||
sort_col = getattr(Company, sort_by, Company.name)
|
||||
if sort_order == "desc":
|
||||
base = base.order_by(desc(sort_col))
|
||||
else:
|
||||
base = base.order_by(asc(sort_col))
|
||||
|
||||
# Count total (without pagination)
|
||||
count_q = select(func.count()).select_from(base.subquery())
|
||||
total_result = await db.execute(count_q)
|
||||
total = total_result.scalar_one()
|
||||
|
||||
# Paginate
|
||||
offset = (page - 1) * page_size
|
||||
paginated = base.offset(offset).limit(page_size)
|
||||
result = await db.execute(paginated)
|
||||
@@ -116,7 +116,6 @@ async def get_company_detail(
|
||||
return None
|
||||
|
||||
data = _company_to_dict(company, include_contacts=True)
|
||||
# Fetch linked contacts
|
||||
contacts_q = (
|
||||
select(Contact, CompanyContact)
|
||||
.join(CompanyContact, CompanyContact.contact_id == Contact.id)
|
||||
@@ -161,6 +160,11 @@ async def create_company(
|
||||
email=data.get("email"),
|
||||
website=data.get("website"),
|
||||
description=data.get("description"),
|
||||
address_street=data.get("address_street"),
|
||||
address_city=data.get("address_city"),
|
||||
address_zip=data.get("address_zip"),
|
||||
address_country=data.get("address_country"),
|
||||
address_state=data.get("address_state"),
|
||||
created_by=user_id,
|
||||
updated_by=user_id,
|
||||
)
|
||||
@@ -198,7 +202,11 @@ async def update_company(
|
||||
return None
|
||||
|
||||
changes: dict[str, Any] = {}
|
||||
for field in ("name", "account_number", "industry", "phone", "email", "website", "description"):
|
||||
all_fields = (
|
||||
"name", "account_number", "industry", "phone", "email",
|
||||
"website", "description",
|
||||
) + ADDRESS_FIELDS
|
||||
for field in all_fields:
|
||||
if field in data and data[field] is not None:
|
||||
old_val = getattr(company, field)
|
||||
changes[field] = {"old": old_val, "new": data[field]}
|
||||
@@ -233,7 +241,6 @@ async def soft_delete_company(
|
||||
company.updated_by = user_id
|
||||
|
||||
if cascade:
|
||||
# Delete N:M links for this company
|
||||
await db.execute(
|
||||
delete(CompanyContact).where(
|
||||
CompanyContact.company_id == company_id,
|
||||
@@ -264,7 +271,6 @@ async def link_contact(
|
||||
is_primary: bool = False,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Link a contact to a company (N:M). Returns link data or None if either not found."""
|
||||
# Verify company exists and is in tenant
|
||||
comp_q = select(Company).where(
|
||||
Company.id == company_id,
|
||||
Company.tenant_id == tenant_id,
|
||||
@@ -274,7 +280,6 @@ async def link_contact(
|
||||
if comp_result.scalar_one_or_none() is None:
|
||||
return None
|
||||
|
||||
# Verify contact exists and is in tenant
|
||||
cont_q = select(Contact).where(
|
||||
Contact.id == contact_id,
|
||||
Contact.tenant_id == tenant_id,
|
||||
@@ -284,7 +289,6 @@ async def link_contact(
|
||||
if cont_result.scalar_one_or_none() is None:
|
||||
return None
|
||||
|
||||
# Check if link already exists
|
||||
existing_q = select(CompanyContact).where(
|
||||
CompanyContact.company_id == company_id,
|
||||
CompanyContact.contact_id == contact_id,
|
||||
@@ -293,17 +297,11 @@ async def link_contact(
|
||||
existing_result = await db.execute(existing_q)
|
||||
existing = existing_result.scalar_one_or_none()
|
||||
if existing is not None:
|
||||
# Update existing link
|
||||
existing.role_at_company = role_at_company
|
||||
existing.is_primary = is_primary
|
||||
await db.flush()
|
||||
await log_audit(
|
||||
db,
|
||||
tenant_id,
|
||||
user_id,
|
||||
"link",
|
||||
"company_contact",
|
||||
existing.id,
|
||||
db, tenant_id, user_id, "link", "company_contact", existing.id,
|
||||
changes={"company_id": str(company_id), "contact_id": str(contact_id)},
|
||||
)
|
||||
return {
|
||||
@@ -324,12 +322,7 @@ async def link_contact(
|
||||
db.add(link)
|
||||
await db.flush()
|
||||
await log_audit(
|
||||
db,
|
||||
tenant_id,
|
||||
user_id,
|
||||
"link",
|
||||
"company_contact",
|
||||
link.id,
|
||||
db, tenant_id, user_id, "link", "company_contact", link.id,
|
||||
changes={"company_id": str(company_id), "contact_id": str(contact_id)},
|
||||
)
|
||||
return {
|
||||
@@ -362,12 +355,7 @@ async def unlink_contact(
|
||||
await db.delete(link)
|
||||
await db.flush()
|
||||
await log_audit(
|
||||
db,
|
||||
tenant_id,
|
||||
user_id,
|
||||
"unlink",
|
||||
"company_contact",
|
||||
link.id,
|
||||
db, tenant_id, user_id, "unlink", "company_contact", link.id,
|
||||
changes={"company_id": str(company_id), "contact_id": str(contact_id)},
|
||||
)
|
||||
return True
|
||||
@@ -403,20 +391,15 @@ async def export_companies_csv(
|
||||
output = io.StringIO()
|
||||
writer = csv.writer(output)
|
||||
writer.writerow(
|
||||
["id", "name", "account_number", "industry", "phone", "email", "website", "description"]
|
||||
["id", "name", "account_number", "industry", "phone", "email", "website", "description",
|
||||
"address_street", "address_city", "address_zip", "address_country", "address_state"]
|
||||
)
|
||||
for c in companies:
|
||||
writer.writerow(
|
||||
[
|
||||
str(c.id),
|
||||
c.name,
|
||||
c.account_number or "",
|
||||
c.industry or "",
|
||||
c.phone or "",
|
||||
c.email or "",
|
||||
c.website or "",
|
||||
c.description or "",
|
||||
]
|
||||
[str(c.id), c.name, c.account_number or "", c.industry or "",
|
||||
c.phone or "", c.email or "", c.website or "", c.description or "",
|
||||
c.address_street or "", c.address_city or "", c.address_zip or "",
|
||||
c.address_country or "", c.address_state or ""]
|
||||
)
|
||||
return output.getvalue()
|
||||
|
||||
@@ -454,28 +437,16 @@ async def export_companies_xlsx(
|
||||
ws = wb.active
|
||||
ws.title = "Companies"
|
||||
headers = [
|
||||
"id",
|
||||
"name",
|
||||
"account_number",
|
||||
"industry",
|
||||
"phone",
|
||||
"email",
|
||||
"website",
|
||||
"description",
|
||||
"id", "name", "account_number", "industry", "phone", "email", "website", "description",
|
||||
"address_street", "address_city", "address_zip", "address_country", "address_state",
|
||||
]
|
||||
ws.append(headers)
|
||||
for c in companies:
|
||||
ws.append(
|
||||
[
|
||||
str(c.id),
|
||||
c.name,
|
||||
c.account_number or "",
|
||||
c.industry or "",
|
||||
c.phone or "",
|
||||
c.email or "",
|
||||
c.website or "",
|
||||
c.description or "",
|
||||
]
|
||||
[str(c.id), c.name, c.account_number or "", c.industry or "",
|
||||
c.phone or "", c.email or "", c.website or "", c.description or "",
|
||||
c.address_street or "", c.address_city or "", c.address_zip or "",
|
||||
c.address_country or "", c.address_state or ""]
|
||||
)
|
||||
|
||||
buf = io.BytesIO()
|
||||
@@ -489,7 +460,6 @@ async def get_company_emails(
|
||||
company_id: uuid.UUID,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Get emails for a company (mail plugin inactive — returns empty array)."""
|
||||
# Verify company exists
|
||||
q = select(Company).where(
|
||||
Company.id == company_id,
|
||||
Company.tenant_id == tenant_id,
|
||||
@@ -497,5 +467,5 @@ async def get_company_emails(
|
||||
)
|
||||
result = await db.execute(q)
|
||||
if result.scalar_one_or_none() is None:
|
||||
return [] # Caller should check existence separately
|
||||
return []
|
||||
return []
|
||||
|
||||
Reference in New Issue
Block a user