"""Context Strategy Experiments — Test all context strategies on RAGSAK. Preserves results in experiments/results/ for future reference. """ import json import os import time from collections import defaultdict, Counter from pathlib import Path from bex.tag_preprocessor.analyze import ( _preprocess_files, _file_to_package, frequency_filter, scan_directory ) from bex.twotinf import build_soa from bex.rwr0 import rwr0 from bex.crx import CRX from bex.idregex import idregex from bex.mdl import lang_size_score BASE = "/home/tobi/Desktop/kesai/RAGSAK" RESULTS_DIR = Path(__file__).parent / "results" def load_data(): """Load all Kotlin files from RAGSAK.""" groups = scan_directory(BASE) kt_files = groups.get(".kt", []) sequences, seq_files = _preprocess_files(kt_files) symbol_seqs = [[text for _, text, _ in seq] for seq in sequences] packages = [_file_to_package(fp, BASE) for fp in seq_files] return symbol_seqs, packages, seq_files def context_file_path_k(symbol_seqs, packages, k): """Option A: Group by last k components of file path.""" contexts = defaultdict(list) for i, seq in enumerate(symbol_seqs): pkg = packages[i] components = pkg.split("/") ctx = tuple(components[-k:]) if len(components) >= k else tuple(components) contexts[ctx].append(seq) return contexts def context_first_k_symbols(symbol_seqs, k): """Option B: Group by first k symbols of call sequence.""" contexts = defaultdict(list) for seq in symbol_seqs: if not seq: contexts[("ε",)].append(seq) else: ctx = tuple(seq[:k]) contexts[ctx].append(seq) return contexts def context_two_d(symbol_seqs, packages, path_k, sym_k): """Option C: Two-dimensional (path_k, first_k_symbols).""" contexts = defaultdict(list) for i, seq in enumerate(symbol_seqs): pkg = packages[i] components = pkg.split("/") path_ctx = tuple(components[-path_k:]) if len(components) >= path_k else tuple(components) if not seq: sym_ctx = ("ε",) else: sym_ctx = tuple(seq[:sym_k]) ctx = path_ctx + sym_ctx contexts[ctx].append(seq) return contexts def context_return_type_heuristic(symbol_seqs): """Option H: Group by return type heuristic (based on last symbol).""" contexts = defaultdict(list) for seq in symbol_seqs: if not seq: contexts[("ε",)].append(seq) else: # Heuristic: last symbol often indicates return type last = seq[-1] if last.startswith("return"): ctx = ("RETURN_" + last.split()[1] if len(last.split()) > 1 else "RETURN_OTHER",) elif last in ("true", "false"): ctx = ("RETURN_BOOL",) elif last.startswith("set") or last.startswith("put"): ctx = ("SIDE_EFFECT",) else: ctx = ("RETURN_VALUE",) contexts[ctx].append(seq) return contexts def evaluate_context(contexts, label, min_methods=3, max_methods=50, max_unique_ratio=0.85, max_alphabet=20, max_soa_edges=100): """Evaluate a context grouping: learn SOREs, collect metrics. Smart filters before attempting REWRITE: - max_methods: skip groups too large (will fail, waste time) - max_unique_ratio: skip if sequences are too diverse - max_alphabet: skip if too many unique symbols (SORE can't compress) - max_soa_edges: skip if SOA is too complex (REWRITE will be slow) """ results = { "strategy": label, "total_contexts": len(contexts), "meaningful_contexts": 0, "total_methods": 0, "methods_in_good_groups": 0, "sore_successes": 0, "sore_failures": 0, "skip_reasons": {}, "groups": [], } for ctx, seqs in sorted(contexts.items(), key=lambda x: -len(x[1])): n = len(seqs) results["total_methods"] += n if n < min_methods: continue results["meaningful_contexts"] += 1 unique = len(set(tuple(s) for s in seqs)) unique_ratio = unique / n if n else 1.0 # Filter: too many methods if n > max_methods: results["skip_reasons"]["too_large"] = results["skip_reasons"].get("too_large", 0) + 1 results["groups"].append({ "context": str(ctx), "methods": n, "unique": unique, "unique_ratio": round(unique_ratio, 3), "sore": "SKIP(too_large)", "sore_success": False, }) continue # Filter: too diverse if unique_ratio > max_unique_ratio: results["skip_reasons"]["too_diverse"] = results["skip_reasons"].get("too_diverse", 0) + 1 results["groups"].append({ "context": str(ctx), "methods": n, "unique": unique, "unique_ratio": round(unique_ratio, 3), "sore": "SKIP(too_diverse)", "sore_success": False, }) continue # Filter: too many unique symbols alphabet = set() for seq in seqs: alphabet.update(seq) if len(alphabet) > max_alphabet: results["skip_reasons"]["large_alphabet"] = results["skip_reasons"].get("large_alphabet", 0) + 1 results["groups"].append({ "context": str(ctx), "methods": n, "unique": unique, "unique_ratio": round(unique_ratio, 3), "sore": "SKIP(large_alphabet)", "sore_success": False, }) continue # Filter: rare symbols filtered = frequency_filter( [[(j, s, 0) for j, s in enumerate(seq)] for seq in seqs], min_coverage=0.0, ) clean = [[text for _, text, _ in r] for r in filtered] clean = [s for s in clean if s] if len(clean) < 2: results["skip_reasons"]["empty_after_filter"] = results["skip_reasons"].get("empty_after_filter", 0) + 1 continue # Build SOA and check complexity before REWRITE soa = build_soa(clean) n_edges = sum(len(v) for v in soa._succ.values()) if n_edges > max_soa_edges: results["skip_reasons"]["complex_soa"] = results["skip_reasons"].get("complex_soa", 0) + 1 results["groups"].append({ "context": str(ctx), "methods": n, "unique": unique, "unique_ratio": round(unique_ratio, 3), "sore": "SKIP(complex_soa)", "sore_success": False, }) continue # Learn SORE sore = rwr0(soa) group_info = { "context": str(ctx), "methods": n, "unique": unique, "unique_ratio": round(unique_ratio, 3), "sore": sore[:200] if sore not in ("∅",) else sore, "sore_success": sore != "∅", } if sore != "∅": results["sore_successes"] += 1 results["methods_in_good_groups"] += n else: results["sore_failures"] += 1 results["groups"].append(group_info) results["coverage"] = ( round(results["methods_in_good_groups"] / results["total_methods"] * 100, 1) if results["total_methods"] > 0 else 0 ) return results def run_experiment(name, contexts, label): """Run one experiment, save results.""" print(f"\n{'='*60}") print(f" {label}") print(f"{'='*60}") t0 = time.time() results = evaluate_context(contexts, label) elapsed = time.time() - t0 results["elapsed_seconds"] = round(elapsed, 2) print(f" Contexts: {results['total_contexts']}") print(f" Meaningful (>=3 methods): {results['meaningful_contexts']}") if results.get("skip_reasons"): for reason, count in results["skip_reasons"].items(): print(f" Skipped ({reason}): {count}") print(f" SORE successes: {results['sore_successes']}") print(f" SORE failures: {results['sore_failures']}") print(f" Coverage: {results['coverage']}%") print(f" Time: {elapsed:.1f}s") # Save out_path = RESULTS_DIR / f"{name}.json" with open(out_path, "w") as f: json.dump(results, f, indent=2) print(f" Saved: {out_path}") return results def main(): print("Loading RAGSAK data...") symbol_seqs, packages, seq_files = load_data() print(f"Loaded {len(symbol_seqs)} methods from {len(set(packages))} packages") all_results = [] # Baseline: package grouping (current approach) contexts = defaultdict(list) for i, seq in enumerate(symbol_seqs): contexts[packages[i]].append(seq) all_results.append(run_experiment("baseline_package", contexts, "Baseline: Package grouping")) # Option A: File path k-equivalence for k in [1, 2, 3]: contexts = context_file_path_k(symbol_seqs, packages, k) all_results.append(run_experiment(f"file_path_k{k}", contexts, f"Option A: File path k={k}")) # Option B: First k symbols for k in [1, 2, 3]: contexts = context_first_k_symbols(symbol_seqs, k) all_results.append(run_experiment(f"first_k_sym_{k}", contexts, f"Option B: First {k} symbols")) # Option C: Two-dimensional for path_k in [1, 2]: for sym_k in [1, 2]: contexts = context_two_d(symbol_seqs, packages, path_k, sym_k) all_results.append(run_experiment( f"two_d_p{path_k}_s{sym_k}", contexts, f"Option C: Path k={path_k} + Symbol k={sym_k}" )) # Option H: Return type heuristic contexts = context_return_type_heuristic(symbol_seqs) all_results.append(run_experiment("return_type_heuristic", contexts, "Option H: Return type heuristic")) # Summary print("\n" + "=" * 80) print(" SUMMARY") print("=" * 80) print(f"{'Strategy':<40} {'Contexts':>8} {'SORE OK':>8} {'Coverage':>8}") print("-" * 80) for r in all_results: print(f"{r['strategy']:<40} {r['meaningful_contexts']:>8} {r['sore_successes']:>8} {r['coverage']:>7}%") # Save summary summary = [ {k: v for k, v in r.items() if k != "groups"} for r in all_results ] summary_path = RESULTS_DIR / "summary.json" with open(summary_path, "w") as f: json.dump(summary, f, indent=2) print(f"\nSummary saved: {summary_path}") if __name__ == "__main__": main()