from datetime import datetime, timezone

from db.models import ProductChannelTaxonomyAssignment, ShopifyCollectionRegistry, ShopifyConnection
from shopify.collection_sync import assign_product_to_collections


class _Client:
    def __init__(self):
        self.calls = []

    def graphql(self, query, variables):
        self.calls.append((query, variables))
        if "query productCollections" in query:
            return {
                "product": {
                    "collections": {
                        "nodes": [
                            {"id": "gid://shopify/Collection/desired"},
                            {"id": "gid://shopify/Collection/stale"},
                        ]
                    }
                }
            }
        if "collectionRemoveProducts" in query:
            return {"collectionRemoveProducts": {"userErrors": []}}
        if "collectionAddProducts" in query:
            return {"collectionAddProducts": {"collection": {"id": variables["id"]}, "userErrors": []}}
        raise AssertionError(query)


def test_canonical_shopify_push_removes_stale_managed_collection(
    catalog_intent_session, monkeypatch
):
    session = catalog_intent_session
    now = datetime.now(timezone.utc)
    session.add(
        ShopifyConnection(
            id=7,
            shop_code="test",
            shop_domain="test.myshopify.com",
            status="active",
        )
    )
    session.flush()
    for collection_id in ("gid://shopify/Collection/desired", "gid://shopify/Collection/stale"):
        session.add(
            ShopifyCollectionRegistry(
                shop_code="test",
                connection_id=7,
                collection_id=collection_id,
                handle=collection_id.rsplit("/", 1)[-1],
                title=collection_id.rsplit("/", 1)[-1],
                collection_type="manual",
                fetched_at=now,
            )
        )
    session.add_all(
        [
            ProductChannelTaxonomyAssignment(
                master_sku="ACH-B12",
                channel_code="shopify",
                connection_id=7,
                taxonomy_kind="collection",
                remote_id="gid://shopify/Collection/desired",
                assignment_status="active",
            ),
            ProductChannelTaxonomyAssignment(
                master_sku="ACH-B12",
                channel_code="shopify",
                connection_id=7,
                taxonomy_kind="collection",
                remote_id="gid://shopify/Collection/stale",
                assignment_status="inactive",
            ),
        ]
    )
    session.flush()
    monkeypatch.setattr(
        "db.product_channel_taxonomy.sku_has_master_taxonomy_placements",
        lambda *_args: True,
    )
    client = _Client()

    applied = assign_product_to_collections(
        session,
        client,
        product_id="gid://shopify/Product/1",
        master_sku="ACH-B12",
        fields={},
        connection_id=7,
    )

    assert applied.collection_ids == ["gid://shopify/Collection/desired"]
    removals = [variables for query, variables in client.calls if "collectionRemoveProducts" in query]
    additions = [variables for query, variables in client.calls if "collectionAddProducts" in query]
    assert removals == [
        {
            "id": "gid://shopify/Collection/stale",
            "productIds": ["gid://shopify/Product/1"],
        }
    ]
    assert additions == []
