import json, math, collections, sys
from rapidfuzz.distance import Levenshtein
from normalize import variants, core, numbers, GENERIC

D = '/var/www/html/peta/storage/app/propertylab-catalogue-match'
rc = json.load(open(f'{D}/realtycheck_schemes.json'))
cat = json.load(open(f'{D}/catalogue_my_projects.json'))

def hav(a, b, c, d):
    a, b, c, d = map(math.radians, (a, b, c, d))
    return 6371000 * 2 * math.asin(math.sqrt(math.sin((c - a) / 2) ** 2 + math.cos(a) * math.cos(c) * math.sin((d - b) / 2) ** 2))

def fnum(x):
    try:
        v = float(x); return v if v != 0 else None
    except (TypeError, ValueError):
        return None

# ---- prepare records
for r in rc:
    r['V'] = [core(v) for v in variants(r['display_name'])]
    r['V'] = [v for v in r['V'] if v] or [variants(r['display_name'])[0]] if variants(r['display_name']) else [[]]
    r['lat'], r['lng'] = fnum(r['latitude']), fnum(r['longitude'])
for p in cat:
    vs = variants(p['project_name'])
    p['V'] = [c for c in (core(v) for v in vs) if c] or vs or [[]]
    p['lat'], p['lng'] = fnum(p['latitude']), fnum(p['longitude'])

# ---- token weights: IDF over both corpora; neighbourhood words and place names weigh less
df = collections.Counter()
for rec in rc + cat:
    for t in set(t for v in rec['V'] for t in v):
        df[t] += 1
N = len(rc) + len(cat)
places = set()
for r in rc:
    for f in ('state', 'district', 'mukim'):
        places.update(core(variants(r.get(f) or '')[0]) if variants(r.get(f) or '') else [])
for p in cat:
    for f in ('state', 'area'):
        places.update(core(variants(p.get(f) or '')[0]) if variants(p.get(f) or '') else [])
def weight(t):
    w = math.log(N / (1 + df[t])) + 0.5
    if t in GENERIC: w *= 0.3
    elif t in places: w *= 0.5
    if any(c.isdigit() for c in t): w = max(w, 1.5)
    return w
W = {}
def wt(t):
    if t not in W: W[t] = weight(t)
    return W[t]

def tok_eq(a, b):
    if a == b: return 1.0
    if a.isdigit() or b.isdigit() or len(a) < 5 or len(b) < 5: return 0.0
    s = Levenshtein.normalized_similarity(a, b)
    return s if s >= 0.84 else 0.0

def soft_jaccard(A, B):
    A, B = list(dict.fromkeys(A)), list(dict.fromkeys(B))
    if not A or not B: return 0.0, 0.0
    matched_w, used = 0.0, set()
    for a in A:
        best, bj = 0.0, None
        for j, b in enumerate(B):
            if j in used: continue
            e = tok_eq(a, b)
            if e > best: best, bj = e, j
        if bj is not None and best > 0:
            used.add(bj); matched_w += best * (wt(a) + wt(B[bj])) / 2
    total = sum(wt(a) for a in A) + sum(wt(b) for b in B) - matched_w
    return matched_w / total if total else 0.0, matched_w

def name_score(rv, pv):
    best = (0.0, 0.0, None, None)
    for a in rv:
        for b in pv:
            s, mw = soft_jaccard(a, b)
            if s > best[0]: best = (s, mw, a, b)
    return best

# ---- candidate blocking: any catalogue record sharing a token (or a near-spelling of one)
post = collections.defaultdict(set)
for i, p in enumerate(cat):
    for v in p['V']:
        for t in v:
            post[t].add(i)
vocab = list(post)
by_len = collections.defaultdict(list)
for t in vocab:
    if len(t) >= 5 and not t.isdigit(): by_len[len(t)].append(t)
fuzzy_cache = {}
def near_tokens(t):
    if t in fuzzy_cache: return fuzzy_cache[t]
    out = {t} if t in post else set()
    if len(t) >= 5 and not t.isdigit():
        for L in range(len(t) - 2, len(t) + 3):
            for u in by_len.get(L, ()):
                if u[0] == t[0] and tok_eq(t, u) > 0: out.add(u)
    fuzzy_cache[t] = out
    return out

HR = {'Condo/Apartment', 'Serviced Apartment', 'Flat', 'Office/SOHO'}
def thresholds(r):
    if r['category'] in HR and r['precision'] != 'road': return 300, 1000, 2500
    return 1500, 3000, 6000

LANDED_T = {'Terrace House', 'Semi-Detached House', 'Detached House', 'Bungalow', 'Cluster House', 'Town House', 'Low-Cost House', 'Land'}
HIGHRISE_T = {'Condominium/Apartment', 'Flat', 'Hotel/Service Apartment', 'Serviced Apartment'}
def type_fit(r, p):
    pt = p.get('property_type') or ''
    parts = {x.strip() for x in pt.split(',') if x.strip()}
    known = {x for x in parts if x in LANDED_T or x in HIGHRISE_T}
    if not known: return 'unknown'
    c = r['category']
    if c == 'Landed': return 'match' if known & LANDED_T else 'conflict'
    if c in ('Condo/Apartment', 'Serviced Apartment', 'Flat'): return 'match' if known & HIGHRISE_T else 'conflict'
    if c == 'Office/SOHO': return 'partial' if known & {'Hotel/Service Apartment', 'Serviced Apartment', 'Condominium/Apartment'} else 'conflict'
    return 'conflict'   # Shop / Industrial: the catalogue is residential

def psf_ratio(r, p):
    a, b = fnum(r.get('reported_psf')), fnum(p.get('psf_median'))
    return round(a / b, 2) if a and b else None

results = []
for n, r in enumerate(rc):
    cand = set()
    for v in r['V']:
        for t in v:
            if wt(t) < 1.0 and len(v) > 1: continue          # do not block on 'taman' alone
            for u in near_tokens(t): cand |= post[u]
    scored = []
    for i in cand:
        p = cat[i]
        s, mw, a, b = name_score(r['V'], p['V'])
        if s < 0.45: continue
        d = hav(r['lat'], r['lng'], p['lat'], p['lng']) if (r['lat'] and p['lat']) else None
        scored.append((s, mw, d, i, a, b))
    strong, ok, far = thresholds(r)
    def rank(x):
        s, mw, d, *_ = x
        dpen = 0 if d is None else (0 if d <= strong else 0.1 if d <= ok else 0.25 if d <= far else 0.6)
        return s - dpen
    scored.sort(key=rank, reverse=True)
    top = []
    for s, mw, d, i, a, b in scored[:3]:
        p = cat[i]
        na, nb = numbers(a), numbers(b)
        top.append({'i': i, 'name_sim': round(s, 3), 'matched_w': round(mw, 2), 'dist_m': None if d is None else round(d),
                    'exact_core': sorted(a) == sorted(b), 'num_a': sorted(na), 'num_b': sorted(nb),
                    'num_conflict': bool(na and nb and na != nb), 'num_one_side': bool(na) != bool(nb),
                    'type_fit': type_fit(r, p), 'psf_ratio': psf_ratio(r, p), 'rank': round(rank((s, mw, d)), 3),
                    'rc_variant': ' '.join(a), 'cat_variant': ' '.join(b)})
    results.append({'scheme_id': r['scheme_id'], 'top': top, 'n_cand': len(scored)})
    if n % 2000 == 0: print(n, file=sys.stderr)

json.dump(results, open(f'{D}/match_raw.json', 'w'))
print('done', len(results))
