IDP-UI / src /Invoice Validation Functions /tax validation.py
Bhuvi13's picture
Upload 20 files
6f0c329 verified
Raw
History Blame Contribute Delete
21.6 kB
from decimal import Decimal, ROUND_HALF_UP
from datetime import date, datetime
def run_tax_amount_validation(
*,
document,
extracted_fields,
line_items,
vendor,
entity,
tax_master_rows,
platform_configs=None,
):
"""
Tax percentage and tax amount validation.
Advisory checks:
- tax rate exists in tax_master
- tax amount matches taxable base
- cross-country routing
- US sales/use tax advisories
- UK VAT advisories
- India GST advisories
- EU reverse charge
- vendor withholding advisory
- vendor tax exemption handling
"""
platform_configs = platform_configs or {}
tax_rate_tolerance = Decimal(
str(platform_configs.get("validation.tax_rate_match_tolerance", "0.0001"))
)
currency_minor_units = platform_configs.get(
"iso_currencies.minor_unit_value",
{
"USD": Decimal("0.01"),
"GBP": Decimal("0.01"),
"EUR": Decimal("0.01"),
"INR": Decimal("0.01"),
"JPY": Decimal("1"),
"BHD": Decimal("0.001"),
},
)
eu_member_states = set(
platform_configs.get(
"validation.eu_member_states",
[
"AT", "BE", "BG", "HR", "CY", "CZ", "DK", "EE", "FI", "FR",
"DE", "GR", "HU", "IE", "IT", "LV", "LT", "LU", "MT", "NL",
"PL", "PT", "RO", "SK", "SI", "ES", "SE",
],
)
)
def as_decimal(value):
if value is None or value == "":
return None
try:
return Decimal(str(value))
except Exception:
return None
def country(value):
if not value:
return None
return str(value).upper()
def is_zero(value):
value = as_decimal(value)
return value is not None and value == Decimal("0")
def is_non_zero(value):
value = as_decimal(value)
return value is not None and value != Decimal("0")
def invoice_date_value():
value = getattr(extracted_fields, "invoice_date", None)
if isinstance(value, datetime):
return value.date()
if isinstance(value, date):
return value
return None # do not invent today's date
def currency_tolerance():
invoice_currency = country(getattr(extracted_fields, "currency", None))
raw = currency_minor_units.get(invoice_currency)
return Decimal(str(raw)) if raw is not None else None
def add_flag(code, *, weight=None, detail=None):
flag = {"code": code, "severity": "advisory"}
if weight is not None:
flag["weight"] = weight
if detail is not None:
flag["detail"] = detail
risk_flags.append(flag)
def add_advisory(code, *, detail=None):
advisory = {"code": code, "severity": "advisory"}
if detail is not None:
advisory["detail"] = detail
# avoid duplicates
if not any(a["code"] == code for a in advisory_flags):
advisory_flags.append(advisory)
def add_skip(step_name, reason):
skipped_steps.append({"step": step_name, "skip_reason": reason})
def effective_tax_master_rows():
invoice_date = invoice_date_value()
rows = []
for row in tax_master_rows:
if getattr(row, "tenant_id", None) != document.tenant_id:
continue
if not getattr(row, "is_active", True):
continue
if getattr(row, "deleted_at", None) is not None:
continue
# If invoice date is unavailable, do not invent one.
# Keep only active rows, but skip date-sensitive matching later if needed.
if invoice_date is not None:
effective_from = getattr(row, "effective_from", None)
effective_to = getattr(row, "effective_to", None)
if effective_from is not None and effective_from > invoice_date:
continue
if effective_to is not None and invoice_date >= effective_to:
continue
rows.append(row)
return rows
def rate_matches(master_rate, invoice_rate):
master_rate = as_decimal(master_rate)
invoice_rate = as_decimal(invoice_rate)
if master_rate is None or invoice_rate is None:
return False
return abs(master_rate - invoice_rate) <= tax_rate_tolerance
def row_region_matches(row_region, hint_region):
if row_region is None:
return True
return hint_region is not None and str(row_region).upper() == str(hint_region).upper()
def find_tax_master_rate(
*,
rate,
country_code,
region_code=None,
allowed_tax_types=None,
reduced_rate_allowed=False,
):
country_code = country(country_code)
matches = []
for row in active_tax_master:
if country(getattr(row, "country_code", None)) != country_code:
continue
if not row_region_matches(getattr(row, "region_code", None), region_code):
continue
if allowed_tax_types is not None and getattr(row, "tax_type", None) not in allowed_tax_types:
continue
if rate_matches(getattr(row, "tax_rate", None), rate):
matches.append(row)
continue
elif reduced_rate_allowed:
tax_name = str(getattr(row, "tax_name", "") or "").lower()
if "reduced" in tax_name:
matches.append(row)
return matches
def find_india_gst_aggregate_rate(rate, region_code, jurisdiction):
"""
Match invoice aggregate GST rate against the sum of applicable component rates:
inter_state : IGST single-row match
intra_state : CGST + SGST pair whose rates sum to the invoice rate
None/unknown: try IGST, then CGST+SGST, then CGST+UTGST
"""
def in_rows(tax_type, rgn=None):
return [
r for r in active_tax_master
if country(getattr(r, "country_code", None)) == "IN"
and getattr(r, "tax_type", None) == tax_type
and row_region_matches(getattr(r, "region_code", None), rgn)
]
def sum_pair(type_a, type_b, rgn=None):
for a in in_rows(type_a, rgn):
a_rate = as_decimal(getattr(a, "tax_rate", None))
if a_rate is None:
continue
for b in in_rows(type_b, rgn):
b_rate = as_decimal(getattr(b, "tax_rate", None))
if b_rate is None:
continue
if abs((a_rate + b_rate) - rate) <= tax_rate_tolerance:
return [a, b]
return []
if jurisdiction == "inter_state":
for r in in_rows("IGST", region_code):
if rate_matches(getattr(r, "tax_rate", None), rate):
return [r]
return []
if jurisdiction == "intra_state":
return sum_pair("CGST", "SGST", region_code)
# Unknown jurisdiction: try all applicable component combinations.
for r in in_rows("IGST"):
if rate_matches(getattr(r, "tax_rate", None), rate):
return [r]
result = sum_pair("CGST", "SGST")
if result:
return result
return sum_pair("CGST", "UTGST")
def all_line_rates_zero():
for line in line_items:
rate = as_decimal(getattr(line, "tax_rate_per_item", None))
if rate is None:
continue
if rate != Decimal("0"):
return False
return True
def sender_tax_family():
return getattr(extracted_fields, "sender_tax_id_inferred_family", None)
def sender_tax_valid_for_country(expected_family):
return (
getattr(extracted_fields, "sender_tax_id", None) is not None
and sender_tax_family() == expected_family
and bool(getattr(extracted_fields, "sender_tax_id_format_valid", True))
and bool(getattr(extracted_fields, "sender_tax_id_checksum_valid", True))
)
def supplier_and_recipient_are_eu():
return supplier_country in eu_member_states and recipient_country in eu_member_states
def eu_reverse_charge_conditions_met():
if not supplier_and_recipient_are_eu():
return False
if supplier_country == recipient_country:
return False
if vendor_exemption_status == "reverse_charge":
return True
supplier_family = f"VAT-{supplier_country}"
supplier_vat_valid = sender_tax_valid_for_country(supplier_family)
recipient_vat_present = bool(getattr(entity, "vat_id", None))
return (
supplier_vat_valid
and recipient_vat_present
and all_line_rates_zero()
and is_zero(getattr(extracted_fields, "tax_amount", None))
)
def india_jurisdiction():
supplier_state = getattr(vendor, "billing_state", None)
recipient_state = getattr(entity, "region_code", None)
if not supplier_state or not recipient_state:
return None
if str(supplier_state).upper() == str(recipient_state).upper():
return "intra_state"
return "inter_state"
risk_flags = []
advisory_flags = []
skipped_steps = []
ops_alerts = []
matched_tax_master_rows = {}
active_tax_master = effective_tax_master_rows()
supplier_country = country(getattr(vendor, "billing_country", None))
recipient_country = country(getattr(entity, "country_code", None))
invoice_currency = country(getattr(extracted_fields, "currency", None))
vendor_exemption_status = getattr(vendor, "exemption_status", None)
invoice_date = invoice_date_value()
if not active_tax_master:
ops_alerts.append(
{
"code": "alert.tax_master_unconfigured",
"tenant_id": document.tenant_id,
}
)
return {
"risk_flags": risk_flags,
"advisory_flags": advisory_flags,
"skipped_steps": [
{"step": "all_tax_checks", "skip_reason": "tax_master_unconfigured"}
],
"ops_alerts": ops_alerts,
"matched_tax_master_rows": matched_tax_master_rows,
}
# Check 1 and Check 2 both skip if supplier country is missing.
if not supplier_country:
add_skip("rate_exists_in_master", "supplier_country_missing")
add_skip("tax_amount_matches_taxable_base", "supplier_country_missing")
else:
unique_rates = set()
# Gather unique (rate, region) pairs — deduplicated by value per SDD §8.1.
for line in line_items:
rate = as_decimal(getattr(line, "tax_rate_per_item", None))
if rate is not None and rate != Decimal("0"):
unique_rates.add(
(
rate,
getattr(line, "region_code_hint", None),
)
)
# Summary tax rate is checked separately
summary_rate = as_decimal(getattr(extracted_fields, "tax_rate", None))
if summary_rate is not None and summary_rate != Decimal("0"):
unique_rates.add((summary_rate, None))
for rate, region_code in unique_rates:
# India aggregate rate matching is handled entirely in Check 6.
if supplier_country == "IN":
continue
# US state tax needs a region hint
if supplier_country == "US" and not region_code:
add_advisory("us_region_unresolved")
continue
allowed_tax_types = None
if supplier_country == "GB":
allowed_tax_types = {"VAT"}
reduced_rate_allowed = vendor_exemption_status == "reduced_rate"
matches = find_tax_master_rate(
rate=rate,
country_code=supplier_country,
region_code=region_code,
allowed_tax_types=allowed_tax_types,
reduced_rate_allowed=reduced_rate_allowed,
)
if matches:
matched_tax_master_rows[str(rate)] = [getattr(row, "id", None) for row in matches]
else:
add_flag(
"tax_rate_not_in_master",
weight=Decimal("0.35"),
detail={
"rate": str(rate),
"country": supplier_country,
"region_code": region_code,
},
)
# Check 2: Amount matches taxable base.
if supplier_country:
tolerance = currency_tolerance()
if tolerance is None:
add_skip("tax_amount_matches_taxable_base", "currency_minor_unit_unknown")
else:
for line in line_items:
rate = as_decimal(getattr(line, "tax_rate_per_item", None))
amount = as_decimal(getattr(line, "amount", None))
actual_tax = as_decimal(getattr(line, "tax_amount_per_item", None))
if rate is None or amount is None or actual_tax is None:
continue
discount = as_decimal(getattr(line, "discount_amount_per_item", None)) or Decimal("0")
taxable_base = amount - discount
expected_tax = taxable_base * (rate / Decimal("100"))
if abs(expected_tax - actual_tax) > tolerance:
add_flag(
"tax_amount_mismatch",
weight=Decimal("0.35"),
detail={
"line_number": getattr(line, "line_number", None),
"expected_tax": str(expected_tax.quantize(tolerance, rounding=ROUND_HALF_UP)),
"actual_tax": str(actual_tax),
"currency": invoice_currency,
},
)
# §8.9: vendor_tax_exempt advisory fires when all extracted line rates are zero.
if vendor_exemption_status == "exempt" and all_line_rates_zero():
add_advisory("vendor_tax_exempt")
# Check 3: Cross-country routing.
summary_tax_amount = as_decimal(getattr(extracted_fields, "tax_amount", None))
if supplier_country is None or summary_tax_amount is None or summary_tax_amount == Decimal("0"):
add_skip("cross_country_routing", "supplier_country_or_tax_amount_missing_or_zero")
elif recipient_country is None:
add_skip("cross_country_routing", "recipient_country_missing")
elif supplier_country != recipient_country:
add_flag(
"cross_country_tax",
weight=Decimal("0.30"),
detail={
"supplier_country": supplier_country,
"recipient_country": recipient_country,
},
)
# Check 4: United States.
if supplier_country == "US" or recipient_country == "US":
if supplier_country == "US" and summary_tax_amount is not None and summary_tax_amount != Decimal("0"):
add_flag(
"us_sales_tax",
weight=Decimal("0.20"),
detail={"tax_amount": str(summary_tax_amount)},
)
if (
recipient_country == "US"
and supplier_country != recipient_country
and all_line_rates_zero()
):
add_flag(
"us_use_tax_possibly_owed",
weight=Decimal("0.15"),
)
# Check 5: United Kingdom.
if supplier_country == "GB" or recipient_country == "GB":
for line in line_items:
rate = as_decimal(getattr(line, "tax_rate_per_item", None))
tax_amount = as_decimal(getattr(line, "tax_amount_per_item", None))
if rate == Decimal("0"):
gb_matches = find_tax_master_rate(
rate=rate,
country_code="GB",
region_code=getattr(line, "region_code_hint", None),
allowed_tax_types={"zero_rated", "exempt", "VAT"},
reduced_rate_allowed=False,
)
if any(getattr(row, "tax_type", None) == "zero_rated" for row in gb_matches):
add_advisory("uk_zero_rated", detail={"line_number": getattr(line, "line_number", None)})
elif tax_amount == Decimal("0") and any(
getattr(row, "tax_type", None) == "exempt" for row in gb_matches
):
add_advisory("uk_exempt", detail={"line_number": getattr(line, "line_number", None)})
elif not (
vendor_exemption_status == "reverse_charge"
and supplier_and_recipient_are_eu()
and supplier_country != recipient_country
) and not (
vendor_exemption_status == "exempt"
and all_line_rates_zero()
):
add_flag(
"tax_rate_not_in_master",
weight=Decimal("0.35"),
detail={
"rate": "0",
"country": "GB",
"source": getattr(line, "line_number", None),
},
)
if (
supplier_country == "GB"
and recipient_country == "GB"
and summary_tax_amount is not None
and summary_tax_amount != Decimal("0")
and not sender_tax_valid_for_country("VAT-GB")
):
add_flag(
"missing_tax_id_for_local_tax",
weight=Decimal("0.25"),
detail={"expected_family": "VAT-GB"},
)
# Check 6: India GST.
if supplier_country == "IN" or recipient_country == "IN":
jurisdiction = india_jurisdiction()
if jurisdiction is None:
add_advisory("in_gst_jurisdiction_unresolved")
else:
add_advisory("in_gst_jurisdiction_inferred", detail={"jurisdiction": jurisdiction})
# Sender tax ID must look like GSTIN-IN or PAN-IN
if sender_tax_family() not in {"GSTIN-IN", "PAN-IN"}:
add_flag(
"missing_tax_id_for_local_tax",
weight=Decimal("0.25"),
detail={"expected_family": "GSTIN-IN or PAN-IN"},
)
for line in line_items:
rate = as_decimal(getattr(line, "tax_rate_per_item", None))
if rate is None or rate == Decimal("0"):
continue
region_code = None
if jurisdiction == "intra_state":
region_code = getattr(vendor, "billing_state", None) or getattr(entity, "region_code", None)
matches = find_india_gst_aggregate_rate(rate, region_code, jurisdiction)
if matches:
matched_tax_master_rows[getattr(line, "line_number", None)] = [
getattr(r, "id", None) for r in matches
]
else:
add_flag(
"tax_rate_not_in_master",
weight=Decimal("0.35"),
detail={
"rate": str(rate),
"country": "IN",
"region_code": region_code,
"source": getattr(line, "line_number", None),
},
)
# Check 7: EU reverse charge.
# If conditions are met, add advisory and retract cross_country_tax raised by Check 3.
if supplier_and_recipient_are_eu() and supplier_country != recipient_country:
if eu_reverse_charge_conditions_met():
add_advisory("eu_reverse_charge")
risk_flags[:] = [f for f in risk_flags if f["code"] != "cross_country_tax"]
# Check 8: Withholding tax advisory.
if vendor_exemption_status == "withholding_applicable":
add_advisory("vendor_withholding_applicable")
for line in line_items:
tax_master_id = getattr(line, "tax_master_id", None)
if not tax_master_id:
continue
for row in active_tax_master:
if getattr(row, "id", None) == tax_master_id and getattr(row, "tax_type", None) == "withholding":
add_advisory("vendor_withholding_applicable")
break
return {
"risk_flags": risk_flags,
"advisory_flags": advisory_flags,
"skipped_steps": skipped_steps,
"ops_alerts": ops_alerts,
"matched_tax_master_rows": matched_tax_master_rows,
}