from __future__ import annotations import json import re from collections import Counter from datetime import datetime, timezone from sqlalchemy import func, select from app.modules.accounting.gstr2b_models import AccountingGSTR2BPurchase from app.modules.accounting.gstr2b_service import analyze_purchase from app.modules.accounting.purchase_enrichment_models import ( AccountingPurchaseEnrichmentBatch, AccountingPurchaseEnrichmentItem, AccountingPurchaseEnrichmentRecord, ) from app.modules.accounting.purchase_enrichment_parser import file_sha256, parse_enrichment def _utcnow(): return datetime.now(timezone.utc) def _norm_invoice(value: str) -> str: return re.sub(r"[^A-Z0-9]", "", str(value or "").upper()) def _value_close(a: float, b: float) -> bool: a = float(a or 0) b = float(b or 0) tolerance = max(2.0, abs(a) * 0.005) return abs(a - b) <= tolerance def _candidate_rows(db, record: AccountingPurchaseEnrichmentRecord): stmt = select(AccountingGSTR2BPurchase).where( AccountingGSTR2BPurchase.tenant_id == record.tenant_id, AccountingGSTR2BPurchase.client_id == record.client_id, ) if record.supplier_gstin: stmt = stmt.where(AccountingGSTR2BPurchase.supplier_gstin == record.supplier_gstin) return list(db.execute(stmt).scalars().all()) def match_record(db, record: AccountingPurchaseEnrichmentRecord): rows = _candidate_rows(db, record) inv = _norm_invoice(record.document_number) date = record.document_date exact = [ r for r in rows if _norm_invoice(r.invoice_number) == inv and r.invoice_date == date and r.document_type == record.document_type ] if len(exact) == 1: record.gstr2b_purchase_id = exact[0].id record.match_status = "linked" record.match_method = "gstin_invoice_date_document_type" record.match_confidence = 100 return exact[0] if len(exact) > 1: record.match_status = "ambiguous" record.match_method = "multiple_exact_candidates" record.match_confidence = 70 record.gstr2b_purchase_id = None return None same_invoice = [r for r in rows if _norm_invoice(r.invoice_number) == inv] if len(same_invoice) == 1: record.gstr2b_purchase_id = same_invoice[0].id record.match_status = "linked" record.match_method = "gstin_invoice_number" record.match_confidence = 94 return same_invoice[0] date_value = [ r for r in rows if r.invoice_date == date and _value_close(r.invoice_value, record.invoice_value) ] if len(date_value) == 1: record.gstr2b_purchase_id = date_value[0].id record.match_status = "linked" record.match_method = "gstin_date_invoice_value" record.match_confidence = 86 return date_value[0] record.gstr2b_purchase_id = None record.match_status = "ambiguous" if len(same_invoice) > 1 or len(date_value) > 1 else "unmatched" record.match_method = "multiple_candidates" if record.match_status == "ambiguous" else "no_match" record.match_confidence = 50 if record.match_status == "ambiguous" else 0 return None def enriched_context(db, purchase_id: int): records = list(db.execute(select(AccountingPurchaseEnrichmentRecord).where( AccountingPurchaseEnrichmentRecord.gstr2b_purchase_id == purchase_id, AccountingPurchaseEnrichmentRecord.match_status == "linked", )).scalars().all()) if not records: return {"description": "", "hsn_code": "", "records": [], "items": []} record_ids = [r.id for r in records] items = list(db.execute(select(AccountingPurchaseEnrichmentItem).where( AccountingPurchaseEnrichmentItem.record_id.in_(record_ids) ).order_by( AccountingPurchaseEnrichmentItem.record_id, AccountingPurchaseEnrichmentItem.line_number, )).scalars().all()) descriptions = [] hsns = [] for item in items: if item.product_name: descriptions.append(item.product_name) if item.description_text and item.description_text not in descriptions: descriptions.append(item.description_text) if item.hsn_code: hsns.append(item.hsn_code) # Only use one HSN as the Phase 6 scalar HSN input when all enriched item lines agree. unique_hsn = sorted(set(hsns)) hsn = unique_hsn[0] if len(unique_hsn) == 1 else "" description = " | ".join(dict.fromkeys(descriptions))[:12000] return { "description": description, "hsn_code": hsn, "records": records, "items": items, "all_hsn_codes": unique_hsn, } def reanalyze_linked_purchase(db, purchase: AccountingGSTR2BPurchase): context = enriched_context(db, purchase.id) original_desc = purchase.description_text or "" original_hsn = purchase.hsn_code or "" # Preserve the GSTR-2B source record. Temporarily enrich the classifier inputs only. merged_desc = " | ".join(x for x in (original_desc, context["description"]) if x) saved_desc = purchase.description_text saved_hsn = purchase.hsn_code try: purchase.description_text = merged_desc[:12000] or None if context["hsn_code"]: purchase.hsn_code = context["hsn_code"] analyze_purchase(db, purchase) finally: purchase.description_text = saved_desc purchase.hsn_code = saved_hsn reasons = [] try: reasons = json.loads(purchase.suggestion_explanation_json or "[]") except Exception: reasons = [] if context["records"]: srcs = sorted({r.source_type.replace("_", " ").title() for r in context["records"]}) reasons.insert(0, f"Enriched with linked {' + '.join(srcs)} source data.") if context["all_hsn_codes"]: reasons.insert(1, f"Enriched item HSN(s): {', '.join(context['all_hsn_codes'][:8])}.") purchase.suggestion_explanation_json = json.dumps(reasons, ensure_ascii=False) db.add(purchase) return purchase def _update_batch_counts(db, batch: AccountingPurchaseEnrichmentBatch): rows = list(db.execute(select(AccountingPurchaseEnrichmentRecord).where( AccountingPurchaseEnrichmentRecord.batch_id == batch.id )).scalars().all()) counter = Counter(r.match_status for r in rows) batch.records_linked = counter.get("linked", 0) batch.records_unmatched = counter.get("unmatched", 0) batch.records_ambiguous = counter.get("ambiguous", 0) batch.status = "matched" batch.completed_at_utc = _utcnow() db.add(batch) def import_enrichment( db, *, tenant_id: int, client_id: int, tally_guid: str, source_type: str, source_period: str, filename: str, content: bytes, user_id: int, ): digest = file_sha256(content) old = db.execute(select(AccountingPurchaseEnrichmentBatch).where( AccountingPurchaseEnrichmentBatch.tenant_id == tenant_id, AccountingPurchaseEnrichmentBatch.client_id == client_id, AccountingPurchaseEnrichmentBatch.source_type == source_type, AccountingPurchaseEnrichmentBatch.file_sha256 == digest, )).scalar_one_or_none() if old: return old, True parsed = parse_enrichment(content, filename, source_type) batch = AccountingPurchaseEnrichmentBatch( tenant_id=tenant_id, client_id=client_id, tally_guid=(tally_guid or "").strip(), source_type=source_type, original_filename=(filename or "")[:260], file_sha256=digest, source_period=(source_period or "")[:20], status="importing", records_read=len(parsed), imported_by_user_id=user_id, ) db.add(batch) db.flush() imported = duplicate = 0 linked_purchases = set() for row, items in parsed: existing = db.execute(select(AccountingPurchaseEnrichmentRecord.id).where( AccountingPurchaseEnrichmentRecord.tenant_id == tenant_id, AccountingPurchaseEnrichmentRecord.client_id == client_id, AccountingPurchaseEnrichmentRecord.source_type == source_type, AccountingPurchaseEnrichmentRecord.source_document_key == row["source_document_key"], )).scalar_one_or_none() if existing: duplicate += 1 continue raw_summary = { "source_type": source_type, "source_document_key": row["source_document_key"], "item_count": len(items), } record = AccountingPurchaseEnrichmentRecord( tenant_id=tenant_id, client_id=client_id, batch_id=batch.id, tally_guid=(tally_guid or "").strip(), raw_summary_json=json.dumps(raw_summary, ensure_ascii=False), **row, ) db.add(record) db.flush() for item in items: db.add(AccountingPurchaseEnrichmentItem(record_id=record.id, **item)) purchase = match_record(db, record) if purchase: linked_purchases.add(purchase.id) imported += 1 batch.records_imported = imported batch.records_duplicate = duplicate db.flush() _update_batch_counts(db, batch) for purchase_id in linked_purchases: purchase = db.get(AccountingGSTR2BPurchase, purchase_id) if purchase: reanalyze_linked_purchase(db, purchase) db.commit() db.refresh(batch) return batch, False def rematch_batch(db, batch: AccountingPurchaseEnrichmentBatch): rows = list(db.execute(select(AccountingPurchaseEnrichmentRecord).where( AccountingPurchaseEnrichmentRecord.batch_id == batch.id )).scalars().all()) linked = set() for record in rows: if record.manually_linked and record.gstr2b_purchase_id: linked.add(record.gstr2b_purchase_id) continue purchase = match_record(db, record) if purchase: linked.add(purchase.id) _update_batch_counts(db, batch) for purchase_id in linked: purchase = db.get(AccountingGSTR2BPurchase, purchase_id) if purchase: reanalyze_linked_purchase(db, purchase) db.commit() return len(rows) def manually_link_record(db, record: AccountingPurchaseEnrichmentRecord, purchase: AccountingGSTR2BPurchase): if record.tenant_id != purchase.tenant_id or record.client_id != purchase.client_id: raise ValueError("The enrichment record and GSTR-2B purchase belong to different client scopes.") record.gstr2b_purchase_id = purchase.id record.match_status = "linked" record.match_method = "manual_user_link" record.match_confidence = 100 record.manually_linked = True db.add(record) batch = db.get(AccountingPurchaseEnrichmentBatch, record.batch_id) if batch: _update_batch_counts(db, batch) reanalyze_linked_purchase(db, purchase) db.commit() return record def unlink_record(db, record: AccountingPurchaseEnrichmentRecord): purchase_id = record.gstr2b_purchase_id record.gstr2b_purchase_id = None record.match_status = "unmatched" record.match_method = "manual_unlink" record.match_confidence = 0 record.manually_linked = False db.add(record) batch = db.get(AccountingPurchaseEnrichmentBatch, record.batch_id) if batch: _update_batch_counts(db, batch) if purchase_id: purchase = db.get(AccountingGSTR2BPurchase, purchase_id) if purchase: reanalyze_linked_purchase(db, purchase) db.commit() return record def batches_for_client(db, tenant_id: int, client_id: int, limit: int = 30): return list(db.execute(select(AccountingPurchaseEnrichmentBatch).where( AccountingPurchaseEnrichmentBatch.tenant_id == tenant_id, AccountingPurchaseEnrichmentBatch.client_id == client_id, ).order_by(AccountingPurchaseEnrichmentBatch.id.desc()).limit(limit)).scalars().all()) def records_for_batch(db, batch_id: int, limit: int = 300): return list(db.execute(select(AccountingPurchaseEnrichmentRecord).where( AccountingPurchaseEnrichmentRecord.batch_id == batch_id ).order_by(AccountingPurchaseEnrichmentRecord.id.desc()).limit(limit)).scalars().all()) def items_for_records(db, record_ids): ids = [int(x) for x in record_ids if x] if not ids: return {} rows = list(db.execute(select(AccountingPurchaseEnrichmentItem).where( AccountingPurchaseEnrichmentItem.record_id.in_(ids) ).order_by( AccountingPurchaseEnrichmentItem.record_id, AccountingPurchaseEnrichmentItem.line_number, )).scalars().all()) result = {} for item in rows: result.setdefault(item.record_id, []).append(item) return result