#!/usr/bin/env python3
"""Test script for the domain term endpoints on Airflow webserver.

Usage:
    # Quick connectivity check
    python scripts/test_domain_terms_endpoint.py --check

    # Create a single domain term
    python scripts/test_domain_terms_endpoint.py create --term-original "感染症" --source-language ja --term-en "infection"

    # Bulk create domain terms
    python scripts/test_domain_terms_endpoint.py bulk --terms '感染症:ja,EBITDA:en,売上:ja' --skip-duplicates

    # Update a domain term translation
    python scripts/test_domain_terms_endpoint.py update TERM_ID --term-en "infection"

    # Rename term_original (triggers cascade to all chunks)
    python scripts/test_domain_terms_endpoint.py update TERM_ID --term-original "伝染病"

    # Update multiple fields at once
    python scripts/test_domain_terms_endpoint.py update TERM_ID --term-en "infection" --domain-category "medical"

    # Delete a domain term (cascades to chunk metadata)
    python scripts/test_domain_terms_endpoint.py delete TERM_ID

    # Use custom host
    python scripts/test_domain_terms_endpoint.py --host http://localhost:8081 create --term-original "test" --source-language en

Requires HMAC_KEY_RING env var (or HMAC_KEY_RING_FILE) to be set.
"""

import argparse
import json
import sys
import uuid

from dotenv import load_dotenv
load_dotenv()

import requests
from textiq_hmac import load_key_ring, sign_context_headers

DEFAULT_TENANT_ID = "28aa3971-48af-4f2a-aa73-9e4f89fbfba3"
DEFAULT_NODE_ID = "22222222-2222-2222-2222-222222222222"
DEFAULT_USER_ID = "33333333-3333-3333-3333-333333333333"
DEFAULT_ROLE = "content_specialist"


def sign_request(tenant_id, node_id, user_id, role):
    """Generate HMAC headers using textiq_hmac package."""
    key_ring = load_key_ring()
    headers = {
        "X-Tenant-ID": tenant_id,
        "X-Node-ID": node_id,
        "X-User-ID": user_id,
        "X-Role": role,
    }
    return sign_context_headers(headers, key_ring)


def _print_response(r):
    """Print response details."""
    print(f"  Status: {r.status_code}")
    print(f"  X-Correlation-ID: {r.headers.get('X-Correlation-ID', 'NOT PRESENT')}")
    print(f"  X-Request-ID: {r.headers.get('X-Request-ID', 'NOT PRESENT')}")
    try:
        body = r.json()
        print(f"  Body: {json.dumps(body, indent=2)}")
    except Exception:
        print(f"  Body: {r.text[:500]}")


def check_endpoint(host):
    """Quick connectivity check — GET should return 405 if route exists."""
    url = f"{host}/api/v1/domain-terms/00000000-0000-0000-0000-000000000000"
    print(f"GET {url}")
    try:
        r = requests.get(url, timeout=5)
        print(f"  Status: {r.status_code}")
        if r.status_code in (405, 401):
            print(f"  -> Route exists ({r.status_code} = expected for unauthenticated GET)")
            return True
        elif r.status_code == 404:
            print("  -> Route NOT found. Plugin may not be loaded.")
            print("     Run: docker exec textiq-airflow-webserver airflow plugins")
            return False
        else:
            print(f"  -> Unexpected status. Body: {r.text[:200]}")
            return False
    except requests.ConnectionError as e:
        print(e)
        print(f"  -> Cannot connect to {host}. Is the webserver running?")
        return False


def test_no_auth(host):
    """PATCH without auth headers — should return 401."""
    url = f"{host}/api/v1/domain-terms/00000000-0000-0000-0000-000000000000"
    print(f"\nPATCH {url} (no auth headers)")
    r = requests.patch(url, json={"term_en": "test"}, timeout=10)
    print(f"  Status: {r.status_code}")
    print(f"  Body: {r.text[:200]}")
    if r.status_code == 401:
        print("  -> PASS: Got 401 as expected")
    else:
        print(f"  -> FAIL: Expected 401, got {r.status_code}")


def create_domain_term(host, payload, tenant_id, node_id, user_id, role):
    """POST / — create a single domain term."""
    url = f"{host}/api/v1/domain-terms/"
    headers = sign_request(tenant_id, node_id, user_id, role)
    headers["X-Correlation-ID"] = str(uuid.uuid4())
    headers["Content-Type"] = "application/json"

    print(f"\nPOST {url}")
    print(f"  Payload: {json.dumps(payload, ensure_ascii=False)}")
    print(f"  Tenant: {tenant_id}  Node: {node_id}")

    r = requests.post(url, headers=headers, json=payload, timeout=30)
    _print_response(r)

    if r.status_code == 201:
        print(f"  -> PASS: Term created (id={r.json().get('term_id')})")
    elif r.status_code == 409:
        print("  -> Duplicate term (409)")
    else:
        print(f"  -> FAIL: Expected 201, got {r.status_code}")


def bulk_create_domain_terms(host, terms, skip_duplicates, tenant_id, node_id, user_id, role):
    """POST /bulk — bulk create domain terms."""
    url = f"{host}/api/v1/domain-terms/bulk"
    headers = sign_request(tenant_id, node_id, user_id, role)
    headers["X-Correlation-ID"] = str(uuid.uuid4())
    headers["Content-Type"] = "application/json"

    payload = {"terms": terms, "skip_duplicates": skip_duplicates}

    print(f"\nPOST {url}")
    print(f"  Terms count: {len(terms)}")
    print(f"  skip_duplicates: {skip_duplicates}")
    print(f"  Tenant: {tenant_id}  Node: {node_id}")

    r = requests.post(url, headers=headers, json=payload, timeout=30)
    _print_response(r)

    if r.status_code == 201:
        body = r.json()
        print(f"  -> PASS: created={body['total_created']} skipped={body['total_skipped']}")
    elif r.status_code == 409:
        print("  -> Duplicate terms found (409)")
    else:
        print(f"  -> FAIL: Expected 201, got {r.status_code}")


def update_domain_term(host, term_id, payload, tenant_id, node_id, user_id, role):
    """PATCH /{term_id}"""
    url = f"{host}/api/v1/domain-terms/{term_id}"
    headers = sign_request(tenant_id, node_id, user_id, role)
    headers["X-Correlation-ID"] = str(uuid.uuid4())
    headers["Content-Type"] = "application/json"

    print(f"\nPATCH {url}")
    print(f"  Payload: {json.dumps(payload, ensure_ascii=False)}")
    print(f"  Tenant: {tenant_id}")
    print(f"  Correlation-ID: {headers['X-Correlation-ID']}")

    r = requests.patch(url, headers=headers, json=payload, timeout=30)
    _print_response(r)

    if r.status_code == 200:
        body = r.json()
        if body.get("cascaded"):
            print("  -> PASS: Domain term updated (with cascade to chunks)")
        else:
            print("  -> PASS: Domain term updated (no cascade)")
    elif r.status_code == 404:
        print("  -> Term not found (404)")
    elif r.status_code == 409:
        print("  -> Conflict: term_original already exists (409)")
    else:
        print(f"  -> FAIL: Expected 200, got {r.status_code}")


def delete_domain_term(host, term_id, tenant_id, node_id, user_id, role):
    """DELETE /{term_id} — delete a domain term (cascades to chunk metadata)."""
    url = f"{host}/api/v1/domain-terms/{term_id}"
    headers = sign_request(tenant_id, node_id, user_id, role)
    headers["X-Correlation-ID"] = str(uuid.uuid4())

    print(f"\nDELETE {url}")
    print(f"  Tenant: {tenant_id}")
    print(f"  Correlation-ID: {headers['X-Correlation-ID']}")

    r = requests.delete(url, headers=headers, timeout=30)
    print(f"  Status: {r.status_code}")
    print(f"  X-Correlation-ID: {r.headers.get('X-Correlation-ID', 'NOT PRESENT')}")
    print(f"  X-Request-ID: {r.headers.get('X-Request-ID', 'NOT PRESENT')}")

    if r.status_code == 204:
        print("  -> PASS: Domain term deleted (204 No Content)")
    elif r.status_code == 404:
        print("  -> Term not found (404)")
        try:
            print(f"  Body: {json.dumps(r.json(), indent=2)}")
        except Exception:
            pass
    else:
        print(f"  -> FAIL: Expected 204, got {r.status_code}")
        try:
            print(f"  Body: {json.dumps(r.json(), indent=2)}")
        except Exception:
            print(f"  Body: {r.text[:500]}")


def main():
    parser = argparse.ArgumentParser(description="Test DE domain term endpoint")
    parser.add_argument("--host", default="http://localhost:8081", help="Airflow webserver URL")
    parser.add_argument("--check", action="store_true", help="Quick connectivity check only")
    parser.add_argument("--tenant-id", default=DEFAULT_TENANT_ID)
    parser.add_argument("--node-id", default=DEFAULT_NODE_ID)
    parser.add_argument("--user-id", default=DEFAULT_USER_ID)
    parser.add_argument("--role", default=DEFAULT_ROLE)

    subparsers = parser.add_subparsers(dest="command")

    sp_create = subparsers.add_parser("create", help="Create a single domain term")
    sp_create.add_argument("--term-original", required=True, help="Original term text")
    sp_create.add_argument("--source-language", required=True, help="Source language code (e.g. ja, en, vi)")
    sp_create.add_argument("--term-ja", default=None, help="Japanese translation")
    sp_create.add_argument("--term-en", default=None, help="English translation")
    sp_create.add_argument("--term-vi", default=None, help="Vietnamese translation")
    sp_create.add_argument("--domain-category", default=None, help="Domain category")
    sp_create.add_argument("--confidence", type=float, default=None, help="Confidence 0..1 (omit → API defaults to 1.0)")

    sp_bulk = subparsers.add_parser("bulk", help="Bulk create domain terms")
    sp_bulk.add_argument("--terms", required=True,
                         help="Comma-separated term:lang[:conf] triples (e.g. 'foo:ja,bar:en:0.8,baz:ja')")
    sp_bulk.add_argument("--skip-duplicates", action="store_true", help="Skip duplicate terms instead of failing")

    sp_delete = subparsers.add_parser("delete", help="Delete a domain term (cascades to chunks)")
    sp_delete.add_argument("term_id", help="Domain term UUID")

    sp_update = subparsers.add_parser("update", help="Update a domain term")
    sp_update.add_argument("term_id", help="Domain term UUID")
    sp_update.add_argument("--term-original", default=None, help="New original term (triggers cascade)")
    sp_update.add_argument("--term-ja", default=None, help="Japanese translation")
    sp_update.add_argument("--term-en", default=None, help="English translation")
    sp_update.add_argument("--term-vi", default=None, help="Vietnamese translation")
    sp_update.add_argument("--domain-category", default=None, help="Domain category")

    args = parser.parse_args()

    print(f"=== Domain Term Endpoint Test ({args.host}) ===\n")

    alive = check_endpoint(args.host)
    if not alive:
        sys.exit(1)

    if args.check:
        test_no_auth(args.host)
        return

    if not args.command:
        print("\nNo command specified. Running auth-only tests.\n")
        test_no_auth(args.host)
        return

    if args.command == "create":
        payload = {
            "term_original": args.term_original,
            "source_language": args.source_language,
        }
        if args.term_ja is not None:
            payload["term_ja"] = args.term_ja
        if args.term_en is not None:
            payload["term_en"] = args.term_en
        if args.term_vi is not None:
            payload["term_vi"] = args.term_vi
        if args.domain_category is not None:
            payload["domain_category"] = args.domain_category
        if args.confidence is not None:
            payload["confidence"] = args.confidence

        create_domain_term(
            args.host, payload,
            args.tenant_id, args.node_id, args.user_id, args.role,
        )

    elif args.command == "bulk":
        terms = []
        for pair in args.terms.split(","):
            pair = pair.strip()
            parts = pair.split(":")
            if len(parts) == 3:
                term_text, lang, conf = parts
                terms.append({
                    "term_original": term_text.strip(),
                    "source_language": lang.strip(),
                    "confidence": float(conf.strip()),
                })
            elif len(parts) == 2:
                term_text, lang = parts
                terms.append({"term_original": term_text.strip(), "source_language": lang.strip()})
            else:
                terms.append({"term_original": pair, "source_language": "ja"})

        bulk_create_domain_terms(
            args.host, terms, args.skip_duplicates,
            args.tenant_id, args.node_id, args.user_id, args.role,
        )

    elif args.command == "delete":
        delete_domain_term(
            args.host, args.term_id,
            args.tenant_id, args.node_id, args.user_id, args.role,
        )

    elif args.command == "update":
        payload = {}
        if args.term_original is not None:
            payload["term_original"] = args.term_original
        if args.term_ja is not None:
            payload["term_ja"] = args.term_ja
        if args.term_en is not None:
            payload["term_en"] = args.term_en
        if args.term_vi is not None:
            payload["term_vi"] = args.term_vi
        if args.domain_category is not None:
            payload["domain_category"] = args.domain_category

        if not payload:
            print("  ERROR: At least one field must be provided (--term-original, --term-ja, --term-en, --term-vi, --domain-category)")
            sys.exit(1)

        update_domain_term(
            args.host, args.term_id, payload,
            args.tenant_id, args.node_id, args.user_id, args.role,
        )


if __name__ == "__main__":
    main()
