"""Verify Milvus data after migration."""
import sys
sys.path.insert(0, '/home/neil/Documents/textiq-doc-extraction')

from pymilvus import Collection, connections
from src.core.config import settings
import os

# Use localhost for local testing
os.environ['MILVUS_HOST'] = 'localhost'

def verify_data():
    """Check Milvus collection data."""
    print("="*80)
    print("MILVUS DATA VERIFICATION")
    print("="*80)

    # Connect
    connections.connect(
        alias="default",
        host="localhost",
        port=19530,
    )

    collection = Collection(settings.milvus_collection)
    collection.load()

    # Get entity count
    num_entities = collection.num_entities
    print(f"\nTotal entities in collection: {num_entities}")

    # Query sample with new fields
    print(f"\nQuerying {min(5, num_entities)} sample records...")

    results = collection.query(
        expr="id > 0",
        output_fields=["id", "filename", "document_type", "domain_category", "domain_terms", "text_content"],
        limit=5,
    )

    for i, entity in enumerate(results):
        print(f"\n{'='*60}")
        print(f"Record {i+1}:")
        print(f"  ID: {entity.get('id')}")
        print(f"  Filename: {entity.get('filename', 'N/A')[:50]}")
        print(f"  Type: {entity.get('document_type', 'N/A')}")
        print(f"  Domain Category: {entity.get('domain_category', 'N/A')}")
        print(f"  Domain Terms: {entity.get('domain_terms', [])}")
        print(f"  Text (first 100 chars): {entity.get('text_content', '')[:100]}...")

    # Count by domain category
    print(f"\n{'='*80}")
    print("DOMAIN CATEGORY DISTRIBUTION")
    print("="*80)

    all_results = collection.query(
        expr="id > 0",
        output_fields=["domain_category"],
        limit=10000,  # Get a large sample
    )

    categories = {}
    for r in all_results:
        cat = r.get('domain_category', 'unknown')
        categories[cat] = categories.get(cat, 0) + 1

    for cat, count in sorted(categories.items(), key=lambda x: x[1], reverse=True):
        print(f"  {cat:20}: {count:5} records")

    print(f"\nNote: Records from migration have default values (category='general', terms=[])")
    print(f"      New documents processed through the DAG will have extracted domain data.")

    connections.disconnect("default")

if __name__ == "__main__":
    try:
        verify_data()
    except Exception as e:
        print(f"Error: {e}")
        import traceback
        traceback.print_exc()
