- supplier/buyer: extract FC tax registration (BT-32/BT-49) into new tax_number field; previously only VA (USt-IdNr.) was read - pflichtfelder: critical error only when neither vat_id nor tax_number is present (EN 16931 BR-CO-26) - ustid: tolerate formatting whitespace inside VAT IDs (e.g. 'DE 140 978 617') - line items: normalize unit price by basis quantity (BT-149), so a price per 100 units no longer inflates line totals by factor 100 - supplier email: also look in DefinedTradeContact (BT-42) - rewrite xrechnung fallback test to assert only our own behavior, independent of factur-x version specifics
401 lines
13 KiB
Python
401 lines
13 KiB
Python
"""Validation functions for ZUGFeRD invoices."""
|
||
|
||
import re
|
||
import time
|
||
from typing import Any
|
||
|
||
from pydantic import ValidationError
|
||
|
||
from src.models import (
|
||
ErrorDetail,
|
||
ValidateRequest,
|
||
ValidationResult,
|
||
XmlData,
|
||
)
|
||
from src.utils import amounts_match
|
||
|
||
|
||
def validate_pflichtfelder(xml_data: XmlData) -> list[ErrorDetail]:
|
||
"""Check required fields are present."""
|
||
errors = []
|
||
|
||
def add_error(field: str, severity: str, message: str | None = None) -> None:
|
||
errors.append(
|
||
ErrorDetail(
|
||
check="pflichtfelder",
|
||
field=field,
|
||
error_code="missing_required",
|
||
message=message or f"Required field '{field}' is missing or empty",
|
||
severity=severity,
|
||
)
|
||
)
|
||
|
||
# Critical fields
|
||
if not xml_data.invoice_number or not xml_data.invoice_number.strip():
|
||
add_error("invoice_number", "critical")
|
||
|
||
if not xml_data.invoice_date or not xml_data.invoice_date.strip():
|
||
add_error("invoice_date", "critical")
|
||
|
||
if not xml_data.supplier.name or not xml_data.supplier.name.strip():
|
||
add_error("supplier.name", "critical")
|
||
|
||
# EN 16931 BR-CO-26: seller identified by VAT ID (BT-31) OR tax number (BT-32)
|
||
has_vat_id = bool(xml_data.supplier.vat_id and xml_data.supplier.vat_id.strip())
|
||
has_tax_number = bool(
|
||
xml_data.supplier.tax_number and xml_data.supplier.tax_number.strip()
|
||
)
|
||
if not has_vat_id and not has_tax_number:
|
||
add_error(
|
||
"supplier.vat_id",
|
||
"critical",
|
||
message=(
|
||
"Supplier must have a VAT ID (supplier.vat_id) "
|
||
"or a tax number (supplier.tax_number)"
|
||
),
|
||
)
|
||
|
||
if not xml_data.buyer.name or not xml_data.buyer.name.strip():
|
||
add_error("buyer.name", "critical")
|
||
|
||
if xml_data.totals.net == 0:
|
||
add_error("totals.net", "critical")
|
||
|
||
if xml_data.totals.gross == 0:
|
||
add_error("totals.gross", "critical")
|
||
|
||
if xml_data.totals.vat_total == 0:
|
||
add_error("totals.vat_total", "critical")
|
||
|
||
# Warning fields
|
||
if xml_data.due_date is not None and not xml_data.due_date.strip():
|
||
add_error("due_date", "warning")
|
||
|
||
if (
|
||
xml_data.payment_terms is not None
|
||
and xml_data.payment_terms.iban is not None
|
||
and not xml_data.payment_terms.iban.strip()
|
||
):
|
||
add_error("payment_terms.iban", "warning")
|
||
|
||
# Line items
|
||
if not xml_data.line_items or len(xml_data.line_items) == 0:
|
||
add_error("line_items", "critical")
|
||
else:
|
||
for idx, item in enumerate(xml_data.line_items):
|
||
field_prefix = f"line_items[{idx}]"
|
||
|
||
if not item.description or not item.description.strip():
|
||
add_error(f"{field_prefix}.description", "critical")
|
||
|
||
if item.quantity == 0:
|
||
add_error(f"{field_prefix}.quantity", "critical")
|
||
|
||
if item.unit_price is None:
|
||
add_error(f"{field_prefix}.unit_price", "critical")
|
||
|
||
if item.line_total is None:
|
||
add_error(f"{field_prefix}.line_total", "critical")
|
||
|
||
if item.vat_rate is None:
|
||
add_error(f"{field_prefix}.vat_rate", "warning")
|
||
|
||
return errors
|
||
|
||
|
||
def validate_betraege(xml_data: XmlData) -> list[ErrorDetail]:
|
||
"""Check amount calculations are correct."""
|
||
errors = []
|
||
|
||
def add_mismatch(field: str, expected: float, actual: float | None) -> None:
|
||
errors.append(
|
||
ErrorDetail(
|
||
check="betraege",
|
||
field=field,
|
||
error_code="calculation_mismatch",
|
||
message=f"Calculation mismatch for '{field}': expected {expected}, got {actual}",
|
||
severity="critical",
|
||
)
|
||
)
|
||
|
||
# Check line_total = quantity × unit_price
|
||
for idx, item in enumerate(xml_data.line_items):
|
||
unit_price = item.unit_price
|
||
line_total = item.line_total
|
||
if unit_price is None or line_total is None:
|
||
continue
|
||
|
||
expected_line_total = item.quantity * unit_price
|
||
if not amounts_match(line_total, expected_line_total):
|
||
add_mismatch(
|
||
f"line_items[{idx}].line_total",
|
||
expected_line_total,
|
||
line_total,
|
||
)
|
||
|
||
# Check totals.net = sum(line_items.line_total)
|
||
line_total_sum = 0.0
|
||
for item in xml_data.line_items:
|
||
if item.line_total is not None:
|
||
line_total_sum += item.line_total
|
||
if not amounts_match(xml_data.totals.net, line_total_sum):
|
||
add_mismatch("totals.net", line_total_sum, xml_data.totals.net)
|
||
|
||
# Check vat_breakdown.amount = base × (rate/100)
|
||
for idx, vat_breakdown in enumerate(xml_data.totals.vat_breakdown):
|
||
expected_amount = vat_breakdown.base * (vat_breakdown.rate / 100)
|
||
if not amounts_match(vat_breakdown.amount, expected_amount):
|
||
add_mismatch(
|
||
f"totals.vat_breakdown[{idx}].amount",
|
||
expected_amount,
|
||
vat_breakdown.amount,
|
||
)
|
||
|
||
# Check totals.vat_total = sum(vat_breakdown.amount)
|
||
vat_breakdown_sum = sum(vb.amount for vb in xml_data.totals.vat_breakdown)
|
||
if not amounts_match(xml_data.totals.vat_total, vat_breakdown_sum):
|
||
add_mismatch("totals.vat_total", vat_breakdown_sum, xml_data.totals.vat_total)
|
||
|
||
# Check totals.gross = totals.net + totals.vat_total
|
||
expected_gross = xml_data.totals.net + xml_data.totals.vat_total
|
||
if not amounts_match(xml_data.totals.gross, expected_gross):
|
||
add_mismatch("totals.gross", expected_gross, xml_data.totals.gross)
|
||
|
||
return errors
|
||
|
||
|
||
def validate_ustid(vat_id: str) -> ErrorDetail | None:
|
||
"""Check VAT ID format (returns None if valid)."""
|
||
if not vat_id or not vat_id.strip():
|
||
return ErrorDetail(
|
||
check="ustid",
|
||
field="vat_id",
|
||
error_code="invalid_format",
|
||
message="VAT ID is empty",
|
||
severity="critical",
|
||
)
|
||
|
||
# Whitespace inside VAT IDs is formatting only (e.g. "DE 140 978 617")
|
||
vat_id = re.sub(r"\s+", "", vat_id)
|
||
|
||
# German VAT ID: DE followed by 9 digits
|
||
if vat_id.startswith("DE"):
|
||
if re.match(r"^DE[0-9]{9}$", vat_id):
|
||
return None
|
||
return ErrorDetail(
|
||
check="ustid",
|
||
field="vat_id",
|
||
error_code="invalid_format",
|
||
message=f"Invalid German VAT ID format: {vat_id}",
|
||
severity="critical",
|
||
)
|
||
|
||
# Austrian VAT ID: ATU followed by 8 digits
|
||
if vat_id.startswith("AT"):
|
||
if re.match(r"^ATU[0-9]{8}$", vat_id):
|
||
return None
|
||
return ErrorDetail(
|
||
check="ustid",
|
||
field="vat_id",
|
||
error_code="invalid_format",
|
||
message=f"Invalid Austrian VAT ID format: {vat_id}",
|
||
severity="critical",
|
||
)
|
||
|
||
# Swiss VAT ID: CHE followed by 9 digits and MWST/TVA/IVA suffix
|
||
if vat_id.startswith("CH"):
|
||
if re.match(r"^CHE[0-9]{9}(MWST|TVA|IVA)$", vat_id):
|
||
return None
|
||
return ErrorDetail(
|
||
check="ustid",
|
||
field="vat_id",
|
||
error_code="invalid_format",
|
||
message=f"Invalid Swiss VAT ID format: {vat_id}",
|
||
severity="critical",
|
||
)
|
||
|
||
return ErrorDetail(
|
||
check="ustid",
|
||
field="vat_id",
|
||
error_code="invalid_format",
|
||
message=f"Unknown country code or invalid VAT ID format: {vat_id}",
|
||
severity="critical",
|
||
)
|
||
|
||
|
||
def validate_pdf_abgleich(xml_data: XmlData, pdf_values: dict) -> list[ErrorDetail]:
|
||
"""Compare XML values to PDF extracted values."""
|
||
errors = []
|
||
|
||
def add_mismatch(field: str, xml_value: Any, pdf_value: Any) -> None:
|
||
errors.append(
|
||
ErrorDetail(
|
||
check="pdf_abgleich",
|
||
field=field,
|
||
error_code="pdf_mismatch",
|
||
message=f"PDF mismatch for '{field}': XML has {xml_value}, PDF has {pdf_value}",
|
||
severity="warning",
|
||
)
|
||
)
|
||
|
||
# Invoice number (exact match)
|
||
if "invoice_number" in pdf_values:
|
||
pdf_invoice = pdf_values["invoice_number"]
|
||
if xml_data.invoice_number != pdf_invoice:
|
||
add_mismatch("invoice_number", xml_data.invoice_number, pdf_invoice)
|
||
|
||
# Totals.gross (within tolerance)
|
||
if "totals.gross" in pdf_values:
|
||
try:
|
||
pdf_gross = float(pdf_values["totals.gross"])
|
||
if not amounts_match(xml_data.totals.gross, pdf_gross):
|
||
add_mismatch("totals.gross", xml_data.totals.gross, pdf_gross)
|
||
except (ValueError, TypeError):
|
||
pass
|
||
|
||
# Totals.net (within tolerance)
|
||
if "totals.net" in pdf_values:
|
||
try:
|
||
pdf_net = float(pdf_values["totals.net"])
|
||
if not amounts_match(xml_data.totals.net, pdf_net):
|
||
add_mismatch("totals.net", xml_data.totals.net, pdf_net)
|
||
except (ValueError, TypeError):
|
||
pass
|
||
|
||
# Totals.vat_total (within tolerance)
|
||
if "totals.vat_total" in pdf_values:
|
||
try:
|
||
pdf_vat = float(pdf_values["totals.vat_total"])
|
||
if not amounts_match(xml_data.totals.vat_total, pdf_vat):
|
||
add_mismatch("totals.vat_total", xml_data.totals.vat_total, pdf_vat)
|
||
except (ValueError, TypeError):
|
||
pass
|
||
|
||
return errors
|
||
|
||
|
||
def validate_invoice(request: ValidateRequest) -> ValidationResult:
|
||
"""Run selected validation checks."""
|
||
start_time = time.time()
|
||
all_errors = []
|
||
all_warnings = []
|
||
|
||
checks_run = 0
|
||
checks_passed = 0
|
||
|
||
if not request.checks:
|
||
return ValidationResult(
|
||
is_valid=True,
|
||
errors=[],
|
||
warnings=[],
|
||
summary={
|
||
"total_checks": 0,
|
||
"checks_passed": 0,
|
||
"checks_failed": 0,
|
||
"critical_errors": 0,
|
||
"warnings": 0,
|
||
},
|
||
validation_time_ms=0,
|
||
)
|
||
|
||
try:
|
||
xml_data = XmlData(**request.xml_data)
|
||
except ValidationError as e:
|
||
# Convert Pydantic validation errors to ValidationResult
|
||
validation_errors: list[ErrorDetail] = []
|
||
for error in e.errors():
|
||
loc = error["loc"]
|
||
field = str(loc[0]) if loc else None
|
||
validation_errors.append(
|
||
ErrorDetail(
|
||
check="schema_validation",
|
||
field=field,
|
||
error_code=error["type"],
|
||
message=error["msg"],
|
||
severity="critical",
|
||
)
|
||
)
|
||
return ValidationResult(
|
||
is_valid=False,
|
||
errors=validation_errors,
|
||
warnings=[],
|
||
summary={
|
||
"total_checks": 1,
|
||
"checks_passed": 0,
|
||
"checks_failed": 1,
|
||
"critical_errors": len(validation_errors),
|
||
"warnings": 0,
|
||
},
|
||
validation_time_ms=int((time.time() - start_time) * 1000),
|
||
)
|
||
|
||
# Run requested checks
|
||
for check_name in request.checks:
|
||
check_errors: list[ErrorDetail] = []
|
||
|
||
if check_name == "pflichtfelder":
|
||
check_errors = validate_pflichtfelder(xml_data)
|
||
checks_run += 1
|
||
elif check_name == "betraege":
|
||
check_errors = validate_betraege(xml_data)
|
||
checks_run += 1
|
||
elif check_name == "ustid":
|
||
# Check supplier VAT ID
|
||
if xml_data.supplier.vat_id:
|
||
error = validate_ustid(xml_data.supplier.vat_id)
|
||
if error:
|
||
check_errors.append(error)
|
||
# Check buyer VAT ID if present
|
||
if xml_data.buyer.vat_id:
|
||
error = validate_ustid(xml_data.buyer.vat_id)
|
||
if error:
|
||
check_errors.append(error)
|
||
checks_run += 1
|
||
elif check_name == "pdf_abgleich":
|
||
if request.pdf_text:
|
||
# For simplicity, try to extract values from PDF text
|
||
pdf_values = {}
|
||
try:
|
||
if "Invoice" in request.pdf_text:
|
||
parts = request.pdf_text.split()
|
||
if len(parts) > 1:
|
||
pdf_values["invoice_number"] = parts[1]
|
||
if "Total:" in request.pdf_text:
|
||
parts = request.pdf_text.split("Total:")
|
||
if len(parts) > 1:
|
||
total_str = parts[1].strip().split()[0]
|
||
pdf_values["totals.gross"] = total_str
|
||
except Exception:
|
||
pass
|
||
check_errors = validate_pdf_abgleich(xml_data, pdf_values)
|
||
checks_run += 1
|
||
|
||
# Separate errors and warnings
|
||
critical_errors = [e for e in check_errors if e.severity == "critical"]
|
||
warnings = [e for e in check_errors if e.severity == "warning"]
|
||
all_errors.extend(critical_errors)
|
||
all_warnings.extend(warnings)
|
||
|
||
if len(critical_errors) == 0:
|
||
checks_passed += 1
|
||
|
||
validation_time_ms = int((time.time() - start_time) * 1000)
|
||
|
||
is_valid = len(all_errors) == 0
|
||
|
||
summary = {
|
||
"total_checks": checks_run,
|
||
"checks_passed": checks_passed,
|
||
"checks_failed": checks_run - checks_passed,
|
||
"critical_errors": len(all_errors),
|
||
"warnings": len(all_warnings),
|
||
}
|
||
|
||
return ValidationResult(
|
||
is_valid=is_valid,
|
||
errors=all_errors,
|
||
warnings=all_warnings,
|
||
summary=summary,
|
||
validation_time_ms=validation_time_ms,
|
||
)
|