"""Exact six-item, 100-rule checkout search. Amounts are integer cents."""
from itertools import combinations

SKUS = ("CF", "CE", "PA", "DE", "TP", "SH")
SHELF = (1900, 800, 600, 2700, 900, 1400)
MEMBER = tuple(p * 90 // 100 for p in SHELF)  # R001; each is an exact cent
ALL = (1 << len(SKUS)) - 1

# (campaign, base cents, increment cents, item condition)
CAMPAIGNS = (
    ("Coffee", 100, 45, lambda m: bool(m & (1 << 0))),
    ("Pantry", 75, 40, lambda m: bool(m & ((1 << 1) | (1 << 2)))),
    ("Laundry", 150, 65, lambda m: bool(m & (1 << 3))),
    ("Hair and oral care", 100, 50, lambda m: bool(m & ((1 << 4) | (1 << 5)))),
    ("Breakfast bundle", 200, 80, lambda m: m & 0b000011 == 0b000011),
    ("Pasta and laundry", 250, 80, lambda m: m & 0b001100 == 0b001100),
    ("Cleaning and care", 300, 85, lambda m: m & 0b011000 == 0b011000),
    ("Mixed basket", 200, 75, lambda m: m.bit_count() >= 2),
    ("Large basket", 300, 100, lambda m: m.bit_count() >= 3),
)
RULES = []
for family, base, increment, condition in CAMPAIGNS:
    for tier in range(11):
        RULES.append((len(RULES) + 2, family, 500 + tier * 500,
                      base + tier * increment, condition))
assert len(RULES) == 99 and RULES[0][0] == 2 and RULES[-1][0] == 100

SUBTOTAL = {
    mask: sum(MEMBER[i] for i in range(6) if mask & (1 << i))
    for mask in range(1, 1 << 6)
}
def discount(rule, receipt):
    _, _, threshold, amount, condition = rule
    return amount if SUBTOTAL[receipt] >= threshold and condition(receipt) else 0

def partitions(remaining):
    """Generate each set partition once by fixing the lowest remaining item."""
    if not remaining:
        yield ()
        return
    first = remaining & -remaining
    group = remaining
    while group:
        if group & first:
            for tail in partitions(remaining ^ group):
                yield (group,) + tail
        group = (group - 1) & remaining

def best_assignment(groups):
    """Maximum-weight matching of single-use coupons to receipts."""
    n = len(groups)
    dp = {0: (0, ())}  # assigned-receipt mask -> (savings, assignments)
    for rule in RULES:
        nxt = dp.copy()  # skip this coupon
        for assigned, (savings, chosen) in dp.items():
            for j, group in enumerate(groups):
                if not (assigned & (1 << j)):
                    d = discount(rule, group)
                    if d:
                        new_mask = assigned | (1 << j)
                        candidate = savings + d
                        if candidate > nxt.get(new_mask, (-1, ()))[0]:
                            nxt[new_mask] = (candidate, chosen + ((j, rule[0]),))
        dp = nxt
    return max(dp.values(), key=lambda x: x[0])

partitions_seen = 0
best = (-1, None, None)
for groups in partitions(ALL):
    partitions_seen += 1
    savings, assignment = best_assignment(groups)
    if savings > best[0]:
        best = (savings, groups, assignment)
assert partitions_seen == 203  # Bell number for six distinct items

manual_groups = ((1 << 2) | (1 << 3) | (1 << 4) | (1 << 5),
                 (1 << 0) | (1 << 1))
manual_coupons = (99, 49)
manual_savings = sum(discount(RULES[rid - 2], group)
                     for rid, group in zip(manual_coupons, manual_groups))
assert manual_savings == 1640

savings, groups, assignment = best
choice = dict(assignment)
assert len(choice) == len(set(choice.values()))
assert sum(groups) == ALL and all(a & b == 0 for a, b in combinations(groups, 2))
assert savings == sum(discount(RULES[choice[j] - 2], group)
                      for j, group in enumerate(groups) if j in choice)
member_total = sum(MEMBER)
assert member_total == 7470
print("Rules:", 1 + len(RULES), "| partitions:", partitions_seen)
print("Member subtotal: $%.2f" % (member_total / 100))
print("Manual: coupon savings $%.2f; payable $%.2f" %
      (manual_savings / 100, (member_total - manual_savings) / 100))
print("Exact: coupon savings $%.2f; payable $%.2f" %
      (savings / 100, (member_total - savings) / 100))
print("Gap: $%.2f (%.2f%% of manual spend)" %
      ((savings - manual_savings) / 100,
       100 * (savings - manual_savings) / (member_total - manual_savings)))
for j, group in enumerate(groups):
    names = "+".join(SKUS[i] for i in range(6) if group & (1 << i))
    rid = choice.get(j)
    d = discount(RULES[rid - 2], group) if rid else 0
    print(names, "subtotal $%.2f" % (SUBTOTAL[group] / 100),
          "coupon", "R%03d" % rid if rid else "none",
          "save $%.2f" % (d / 100),
          "pay $%.2f" % ((SUBTOTAL[group] - d) / 100))

# Independent upper bound: permit coupon reuse across receipts.
best_on_receipt = {m: max(discount(r, m) for r in RULES) for m in SUBTOTAL}
relaxed = {0: 0}
for remaining in range(1, ALL + 1):
    first = remaining & -remaining
    upper = 0
    group = remaining
    while group:
        if group & first:
            upper = max(upper, best_on_receipt[group] + relaxed[remaining ^ group])
        group = (group - 1) & remaining
    relaxed[remaining] = upper
print("Relaxed maximum coupon savings: $%.2f" % (relaxed[ALL] / 100))
assert relaxed[ALL] == savings
for j, group in enumerate(groups):
    rid = choice[j]
    rule = RULES[rid - 2]
    assert rule[0] == rid and rule[2] <= SUBTOTAL[group] and rule[4](group)
