import asyncio
from datetime import date, datetime
from decimal import Decimal
import os
import sys
import unittest
from pathlib import Path
from unittest.mock import patch

from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
os.environ.setdefault("DB_NAME", ":memory:.db")
os.environ.setdefault("JWT_SECRET", "test-jwt-secret-for-ai-query-executor-tests")

from app.core.database import Base
from app.core.models import (
    Customer,
    Expense,
    GstReturn,
    Invoice,
    InvoiceStatus,
    Payment,
    Product,
    Reminder,
    ReminderType,
    Subscription,
    SubscriptionStatus,
    Tenant,
)
from app.api.ai import handle_ai_chat_message
from app.services.ai_query_executor import build_query_from_message, execute_ai_query


def _db_session():
    engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False})
    Base.metadata.create_all(engine)
    SessionLocal = sessionmaker(bind=engine)
    db = SessionLocal()
    db.add(Tenant(id=1, company_code="test", name="Test"))
    db.commit()
    return db


def _seed_revenue_records(db):
    customer = Customer(tenant_id=1, name="Acme Ltd", email="billing@acme.test")
    other_customer = Customer(tenant_id=1, name="Outside Co", email="accounts@outside.test")
    db.add_all([customer, other_customer])
    db.flush()

    paid_invoice = Invoice(
        tenant_id=1,
        invoice_number="INV-AI-001",
        customer_id=customer.id,
        invoice_date=date(2026, 6, 2),
        due_date=date(2026, 6, 9),
        total_amount=Decimal("230.00"),
        currency="NZD",
        status=InvoiceStatus.paid,
    )
    unpaid_invoice = Invoice(
        tenant_id=1,
        invoice_number="INV-AI-002",
        customer_id=customer.id,
        invoice_date=date(2026, 6, 3),
        due_date=date(2026, 6, 10),
        total_amount=Decimal("460.00"),
        currency="NZD",
        status=InvoiceStatus.sent,
    )
    old_paid_invoice = Invoice(
        tenant_id=1,
        invoice_number="INV-AI-OLD",
        customer_id=other_customer.id,
        invoice_date=date(2026, 4, 30),
        due_date=date(2026, 5, 7),
        total_amount=Decimal("999.00"),
        currency="NZD",
        status=InvoiceStatus.paid,
    )
    db.add_all([paid_invoice, unpaid_invoice, old_paid_invoice])
    db.flush()

    db.add_all(
        [
            Payment(
                tenant_id=1,
                invoice_id=paid_invoice.id,
                received_date=date(2026, 6, 4),
                paid_at=datetime(2026, 6, 4, 10, 0, 0),
                amount=Decimal("100.00"),
                currency="NZD",
                method="bank_transfer",
                reference="REF-100",
            ),
            Payment(
                tenant_id=1,
                invoice_id=unpaid_invoice.id,
                received_date=date(2026, 6, 5),
                paid_at=datetime(2026, 6, 5, 10, 0, 0),
                amount=Decimal("50.00"),
                currency="NZD",
                method="card",
                reference="REF-50",
            ),
            Payment(
                tenant_id=1,
                invoice_id=old_paid_invoice.id,
                received_date=date(2026, 5, 1),
                paid_at=datetime(2026, 5, 1, 10, 0, 0),
                amount=Decimal("999.00"),
                currency="NZD",
                method="bank_transfer",
                reference="REF-OLD",
            ),
        ]
    )
    db.commit()


def _seed_core_query_records(db):
    customer = Customer(tenant_id=1, name="Beta Ltd", company_name="Beta Trading", email="owner@beta.test", customer_type="project")
    product = Product(tenant_id=1, name="Support Plan", unit_price=Decimal("120.00"), product_type="subscription")
    db.add_all([customer, product])
    db.flush()

    invoice = Invoice(
        tenant_id=1,
        invoice_number="INV-CORE-001",
        customer_id=customer.id,
        invoice_date=date(2026, 6, 7),
        due_date=date(2026, 6, 14),
        total_amount=Decimal("345.00"),
        currency="NZD",
        status=InvoiceStatus.sent,
    )
    expense = Expense(
        tenant_id=1,
        expense_date=date(2026, 6, 8),
        vendor_name="Office Store",
        category="office",
        amount_gross=Decimal("57.50"),
        gst_amount=Decimal("7.50"),
        currency="NZD",
        status="confirmed",
        source="manual",
    )
    subscription = Subscription(
        tenant_id=1,
        customer_id=customer.id,
        product_id=product.id,
        start_date=date(2026, 5, 1),
        end_date=date(2026, 6, 20),
        status=SubscriptionStatus.active,
        auto_renew=True,
        next_invoice_date=date(2026, 6, 18),
    )
    reminder = Reminder(
        tenant_id=1,
        reminder_type=ReminderType.invoice_overdue,
        trigger_days=7,
        send_email=True,
        send_telegram=False,
        status="active",
    )
    gst_return = GstReturn(
        tenant_id=1,
        period_start=date(2026, 4, 1),
        period_end=date(2026, 6, 30),
        status="draft",
        gst_output=Decimal("15.00"),
        gst_input=Decimal("7.50"),
        gst_payable=Decimal("7.50"),
        total_sales_income=Decimal("115.00"),
        total_purchases_expenses=Decimal("57.50"),
    )
    db.add_all([invoice, expense, subscription, reminder, gst_return])
    db.commit()


class AIQueryExecutorTests(unittest.TestCase):
    def test_chinese_current_year_income_message_builds_revenue_query(self):
        query = build_query_from_message("今年的收入是多少？")

        self.assertEqual("query", query["action"])
        self.assertEqual("revenue", query["query_type"])
        self.assertEqual("zh", query["language"])
        self.assertEqual(date.today().replace(month=1, day=1).isoformat(), query["filters"]["date_from"])
        self.assertEqual(date.today().isoformat(), query["filters"]["date_to"])

    def test_english_current_year_income_message_builds_revenue_query(self):
        query = build_query_from_message("What is this year's income?")

        self.assertEqual("query", query["action"])
        self.assertEqual("revenue", query["query_type"])
        self.assertEqual("en", query["language"])
        self.assertEqual(date.today().replace(month=1, day=1).isoformat(), query["filters"]["date_from"])
        self.assertEqual(date.today().isoformat(), query["filters"]["date_to"])

    def test_chinese_named_month_income_message_builds_revenue_query(self):
        query = build_query_from_message("看看五月的收入")
        year = date.today().year

        self.assertEqual("query", query["action"])
        self.assertEqual("revenue", query["query_type"])
        self.assertEqual("zh", query["language"])
        self.assertEqual(date(year, 5, 1).isoformat(), query["filters"]["date_from"])
        self.assertEqual(date(year, 5, 31).isoformat(), query["filters"]["date_to"])

    def test_chat_handles_chinese_income_query_without_calling_ai_service(self):
        db = _db_session()
        try:
            with patch("app.api.ai.chat_with_ai", side_effect=AssertionError("AI service should not be called")):
                response = asyncio.run(handle_ai_chat_message(db, "今年的收入是多少？", tenant_id=1))

            self.assertEqual("query_result", response.action)
            self.assertEqual("database", response.source)
            self.assertEqual("query_revenue", response.tool_name)
            self.assertIn("以下数据来自数据库查询", response.reply)
            self.assertEqual(date.today().replace(month=1, day=1).isoformat(), response.result["period"]["date_from"])
            self.assertEqual(date.today().isoformat(), response.result["period"]["date_to"])
        finally:
            db.close()

    def test_revenue_query_returns_cash_and_paid_invoice_views(self):
        db = _db_session()
        try:
            _seed_revenue_records(db)

            result = execute_ai_query(
                db,
                {
                    "action": "query",
                    "query_type": "revenue",
                    "filters": {
                        "date_from": "2026-06-01",
                        "date_to": "2026-06-30",
                    },
                },
                tenant_id=1,
            )

            self.assertEqual("query_result", result["action"])
            self.assertEqual("database", result["source"])
            self.assertEqual("query_revenue", result["tool_name"])
            self.assertEqual(150.0, result["result"]["cash_received"]["total"])
            self.assertEqual(2, result["result"]["cash_received"]["payment_count"])
            self.assertEqual(230.0, result["result"]["paid_invoices"]["total"])
            self.assertEqual(1, result["result"]["paid_invoices"]["invoice_count"])
            self.assertEqual(
                ["INV-AI-001", "INV-AI-002"],
                [row["invoice_number"] for row in result["result"]["cash_received"]["rows"]],
            )
            self.assertEqual(
                ["INV-AI-001"],
                [row["invoice_number"] for row in result["result"]["paid_invoices"]["rows"]],
            )
            self.assertIn("The following figures come from database queries", result["reply"])
            self.assertIn("INV-AI-001", result["reply"])
            self.assertIn("INV-AI-002", result["reply"])
            self.assertNotIn("INV-AI-OLD", result["reply"])
        finally:
            db.close()

    def test_revenue_query_ignores_model_provided_result_data(self):
        db = _db_session()
        try:
            _seed_revenue_records(db)

            result = execute_ai_query(
                db,
                {
                    "action": "query",
                    "query_type": "stats",
                    "filters": {
                        "date_from": "2026-06-01",
                        "date_to": "2026-06-30",
                        "status": "paid",
                    },
                    "result": {
                        "total_revenue": 47850.0,
                        "invoice_count": 12,
                        "top_customers": [{"name": "Fake Customer", "amount": 47850.0}],
                    },
                },
                tenant_id=1,
            )

            self.assertEqual("database", result["source"])
            self.assertEqual(150.0, result["result"]["cash_received"]["total"])
            self.assertEqual(230.0, result["result"]["paid_invoices"]["total"])
            self.assertNotIn("47850", result["reply"])
            self.assertNotIn("Fake Customer", result["reply"])
        finally:
            db.close()

    def test_chinese_revenue_query_returns_chinese_database_reply(self):
        db = _db_session()
        try:
            _seed_revenue_records(db)

            result = execute_ai_query(
                db,
                {
                    "action": "query",
                    "query_type": "revenue",
                    "language": "zh",
                    "filters": {
                        "date_from": "2026-06-01",
                        "date_to": "2026-06-30",
                    },
                },
                tenant_id=1,
            )

            self.assertEqual("query_result", result["action"])
            self.assertEqual("database", result["source"])
            self.assertIn("以下数据来自数据库查询", result["reply"])
            self.assertIn("实际收款", result["reply"])
            self.assertIn("已付款发票合计", result["reply"])
            self.assertIn("INV-AI-001", result["reply"])
        finally:
            db.close()

    def test_query_requires_supported_type_and_date_range(self):
        db = _db_session()
        try:
            unsupported = execute_ai_query(db, {"action": "query", "query_type": "payroll", "filters": {}}, tenant_id=1)
            self.assertEqual("error", unsupported["action"])
            self.assertIn("cannot run that query yet", unsupported["reply"])

            missing_dates = execute_ai_query(db, {"action": "query", "query_type": "revenue", "filters": {}}, tenant_id=1)
            self.assertEqual("clarify", missing_dates["action"])
            self.assertIn("Please provide a date range", missing_dates["reply"])
        finally:
            db.close()

    def test_query_requires_explicit_tenant_id(self):
        db = _db_session()
        try:
            with self.assertRaises(TypeError):
                execute_ai_query(db, {"action": "query", "query_type": "customer", "filters": {}})
        finally:
            db.close()

    def test_core_query_tools_return_database_rows(self):
        db = _db_session()
        try:
            _seed_core_query_records(db)

            cases = [
                ("invoice", "query_invoices", "INV-CORE-001"),
                ("customer", "query_customers", "Beta Ltd"),
                ("product", "query_products", "Support Plan"),
                ("expense", "query_expenses", "Office Store"),
                ("subscription", "query_subscriptions", "Support Plan"),
                ("reminder", "query_reminders", "invoice_overdue"),
                ("gst", "query_gst_summary", "gst_payable"),
            ]

            for query_type, tool_name, expected_text in cases:
                with self.subTest(query_type=query_type):
                    result = execute_ai_query(
                        db,
                        {
                            "action": "query",
                            "query_type": query_type,
                            "filters": {
                                "date_from": "2026-06-01",
                                "date_to": "2026-06-30",
                            },
                        },
                        tenant_id=1,
                    )

                    self.assertEqual("query_result", result["action"])
                    self.assertEqual("database", result["source"])
                    self.assertEqual(tool_name, result["tool_name"])
                    self.assertGreaterEqual(result["result"]["count"], 1)
                    self.assertIn(expected_text, result["reply"])
        finally:
            db.close()


if __name__ == "__main__":
    unittest.main()
