from db.location_availability_normalize import normalize_location_availability_codes
from db.models import MasterLocationRegistry, MasterProduct, MasterProductLocationAvailability


def test_normalize_location_availability_codes_updates_legacy_keys(catalog_intent_session):
    session = catalog_intent_session
    session.add_all(
        [
            MasterLocationRegistry(code="1", name="Keyport Warehouse", location_kind="warehouse", is_active=True),
            MasterLocationRegistry(code="12", name="Baltimore Warehouse", location_kind="warehouse", is_active=True),
        ]
    )
    product = MasterProduct(sku="ACH-B12", name="ACH-B12", row_hash="h", is_active=True)
    session.add(product)
    session.flush()
    session.add_all(
        [
            MasterProductLocationAvailability(
                product_id=product.id,
                sku=product.sku,
                location_code="location_keyport",
                is_available=True,
                source_value="Yes",
                source_label="test",
            ),
            MasterProductLocationAvailability(
                product_id=product.id,
                sku=product.sku,
                location_code="location_baltimore",
                is_available=False,
                source_value="No",
                source_label="test",
            ),
        ]
    )
    session.commit()

    result = normalize_location_availability_codes(session, dry_run=False)
    session.commit()

    assert result["updated"] == 2
    rows = session.query(MasterProductLocationAvailability).filter_by(product_id=product.id).all()
    assert sorted(row.location_code for row in rows) == ["1", "12"]
