code-shrish/itr-detection
0
1"""2Baseline Inference Script for ITR Fraud Detection Environment.3 4Uses the Google Gemini API to run a model against all 3 tasks.5Produces reproducible baseline scores.6 7Usage:8 export GEMINI_API_KEY="your-key"9 python baseline.py10 11Or run locally without API:12 python baseline.py --local13"""14 15import argparse16import json17import os18import sys19from typing import Any, Dict, List20 21# Attempt to load from .env file if python-dotenv is installed22try:23 from dotenv import load_dotenv # type: ignore24 load_dotenv()25except ImportError:26 pass27 28sys.path.insert(0, os.path.dirname(__file__))29 30from models import ActionType, ITRAction, DocumentType, Severity, VerdictType # type: ignore31from server.itr_environment import ITREnvironment # type: ignore32 33 34def run_openai_agent(env: ITREnvironment, task_id: str) -> Dict[str, Any]:35 """Run OpenAI-powered agent on a task."""36 try:37 from openai import OpenAI38 except ImportError:39 print("⚠️ openai package not installed. Run: pip install openai")40 return run_heuristic_agent(env, task_id)41 42 api_key = os.environ.get("HF_TOKEN", os.environ.get("OPENAI_API_KEY"))43 api_base_url = os.environ.get("API_BASE_URL") 44 model_name = os.environ.get("MODEL_NAME", "gpt-4o-mini")45 46 if not api_key:47 print("⚠️ HF_TOKEN / OPENAI_API_KEY not set. Falling back to heuristic agent.")48 return run_heuristic_agent(env, task_id)49 50 client = OpenAI(base_url=api_base_url, api_key=api_key)51 obs = env.reset(task_id=task_id)52 53 system_prompt = """You are an expert income tax auditor for the Indian Income Tax Department.54You are reviewing an ITR (Income Tax Return) filing for potential fraud.55 56You MUST respond with a valid JSON action. Available action types:571. investigate_field - Investigate a specific field. Params: field_name (e.g., "income.salary", "deductions.section_80c")582. cross_reference - Cross-check two fields. Params: field_a, field_b (e.g., "salary", "tds")593. request_document - Request a document. Params: document_type (form_16, bank_statement, rent_receipts, investment_proofs, capital_gains_statement, business_books, related_party_records)604. flag_anomaly - Flag a suspicious field. Params: anomaly_field, anomaly_reason, anomaly_severity (low/medium/high/critical)615. render_verdict - Final decision. Params: verdict (legitimate/suspicious/fraudulent), confidence (0-1), explanation62 63Respond ONLY with JSON. Example:64{"action_type": "investigate_field", "field_name": "income.salary"}65{"action_type": "flag_anomaly", "anomaly_field": "deductions.section_80c", "anomaly_reason": "Exceeds 1.5L limit", "anomaly_severity": "high"}66{"action_type": "render_verdict", "verdict": "fraudulent", "confidence": 0.9, "explanation": "Multiple anomalies found"}67"""68 69 steps = []70 done = False71 total_reward = 0.072 73 while not done:74 user_msg = json.dumps({75 "step": obs.step_number,76 "max_steps": obs.max_steps,77 "task": obs.task_description,78 "itr_data": obs.itr_summary,79 "last_result": obs.last_action_result,80 "flagged_so_far": obs.flagged_anomalies,81 "investigations": [r.model_dump() for r in obs.investigation_results[-3:]],82 }, indent=2, default=str)83 84 response = client.chat.completions.create(85 model=model_name,86 messages=[87 {"role": "system", "content": system_prompt},88 {"role": "user", "content": user_msg},89 ],90 temperature=0.1,91 max_tokens=500,92 )93 94 action_text = response.choices[0].message.content.strip()95 # Parse JSON from response96 try:97 if "```" in action_text:98 action_text = action_text.split("```")[1]99 if action_text.startswith("json"):100 action_text = action_text[4:]101 action_data = json.loads(action_text)102 except json.JSONDecodeError:103 # Fallback: render verdict if can't parse104 action_data = {105 "action_type": "render_verdict",106 "verdict": "suspicious",107 "confidence": 0.5,108 "explanation": "Unable to determine",109 }110 111 # Build action112 action = ITRAction(**action_data)113 result = env.step(action)114 obs = result.observation115 done = result.done116 total_reward += result.reward117 118 print(f"[STEP] action={action.action_type.value} reward={result.reward:.4f}")119 120 steps.append({121 "step": obs.step_number,122 "action": action_data,123 "reward": result.reward,124 "done": done,125 })126 127 return {128 "task_id": task_id,129 "steps": steps,130 "total_reward": total_reward,131 "final_score": result.info.get("final_score", 0.0),132 "num_steps": len(steps),133 }134 135 136def run_heuristic_agent(env: ITREnvironment, task_id: str) -> Dict[str, Any]:137 """138 Rule-based baseline agent that follows a systematic audit procedure.139 This provides reproducible scores without requiring an API key.140 """141 obs = env.reset(task_id=task_id)142 done = False143 total_reward = 0.0144 steps = []145 146 # Step 1: Investigate salary147 action = ITRAction(148 action_type=ActionType.INVESTIGATE_FIELD,149 field_name="income.salary",150 )151 result = env.step(action)152 obs = result.observation153 done = result.done154 total_reward += result.reward155 steps.append({"action": "investigate income.salary", "reward": result.reward})156 157 if done:158 return _build_result(task_id, steps, total_reward, result)159 160 # Step 2: Cross-reference salary with TDS161 action = ITRAction(162 action_type=ActionType.CROSS_REFERENCE,163 field_a="salary",164 field_b="tds",165 )166 result = env.step(action)167 obs = result.observation168 done = result.done169 total_reward += result.reward170 steps.append({"action": "cross_reference salary↔tds", "reward": result.reward})171 172 if done:173 return _build_result(task_id, steps, total_reward, result)174 175 # Step 3: Request Form 16176 action = ITRAction(177 action_type=ActionType.REQUEST_DOCUMENT,178 document_type=DocumentType.FORM_16,179 )180 result = env.step(action)181 obs = result.observation182 done = result.done183 total_reward += result.reward184 steps.append({"action": "request Form 16", "reward": result.reward})185 186 if done:187 return _build_result(task_id, steps, total_reward, result)188 189 # Step 4: Investigate deductions190 action = ITRAction(191 action_type=ActionType.INVESTIGATE_FIELD,192 field_name="deductions.section_80c",193 )194 result = env.step(action)195 obs = result.observation196 done = result.done197 total_reward += result.reward198 steps.append({"action": "investigate deductions.section_80c", "reward": result.reward})199 200 if done:201 return _build_result(task_id, steps, total_reward, result)202 203 # Step 5: Check for suspicious findings — flag anomalies based on observations204 for inv in obs.investigation_results:205 if inv.suspicious and not done:206 action = ITRAction(207 action_type=ActionType.FLAG_ANOMALY,208 anomaly_field=inv.finding.split(":")[0] if ":" in inv.finding else "detected_field",209 anomaly_reason=inv.detail[:200],210 anomaly_severity=Severity.HIGH,211 )212 result = env.step(action)213 obs = result.observation214 done = result.done215 total_reward += result.reward # type: ignore216 steps.append({"action": f"flag_anomaly: {inv.finding[:50]}", "reward": result.reward})217 if done:218 return _build_result(task_id, steps, total_reward, result)219 220 # Check document discrepancies221 for doc in obs.document_results:222 for disc in doc.discrepancies:223 if not done:224 action = ITRAction(225 action_type=ActionType.FLAG_ANOMALY,226 anomaly_field=doc.document_type,227 anomaly_reason=disc[:200],228 anomaly_severity=Severity.HIGH,229 )230 result = env.step(action)231 obs = result.observation232 done = result.done233 total_reward += result.reward # type: ignore234 steps.append({"action": f"flag_anomaly from doc: {disc[:40]}", "reward": result.reward})235 if done:236 return _build_result(task_id, steps, total_reward, result)237 238 # Additional investigations for medium/hard tasks239 if task_id in ("medium", "hard") and not done:240 for field in ["deductions.hra_exemption", "income.capital_gains_long", "high_value_transactions", "previous_years"]:241 if not done:242 action = ITRAction(243 action_type=ActionType.INVESTIGATE_FIELD,244 field_name=field,245 )246 result = env.step(action)247 obs = result.observation248 done = result.done249 total_reward += result.reward # type: ignore250 steps.append({"action": f"investigate {field}", "reward": result.reward})251 252 if result.observation.investigation_results and result.observation.investigation_results[-1].suspicious and not done:253 inv = result.observation.investigation_results[-1]254 action = ITRAction(255 action_type=ActionType.FLAG_ANOMALY,256 anomaly_field=field,257 anomaly_reason=inv.detail[:200],258 anomaly_severity=Severity.HIGH,259 )260 result = env.step(action)261 obs = result.observation262 done = result.done263 total_reward += result.reward # type: ignore264 steps.append({"action": f"flag {field}", "reward": result.reward})265 if done:266 return _build_result(task_id, steps, total_reward, result)267 268 # Hard task: additional document requests269 if task_id == "hard" and not done:270 for doc_type in [DocumentType.BANK_STATEMENT, DocumentType.BUSINESS_BOOKS, DocumentType.RELATED_PARTY_RECORDS]:271 if not done:272 action = ITRAction(273 action_type=ActionType.REQUEST_DOCUMENT,274 document_type=doc_type,275 )276 result = env.step(action)277 obs = result.observation278 done = result.done279 total_reward += result.reward # type: ignore280 steps.append({"action": f"request {doc_type.value}", "reward": result.reward}) # type: ignore281 282 if result.observation.document_results and result.observation.document_results[-1].discrepancies and not done:283 doc = result.observation.document_results[-1]284 for disc in doc.discrepancies:285 if not done:286 action = ITRAction(287 action_type=ActionType.FLAG_ANOMALY,288 anomaly_field=doc.document_type,289 anomaly_reason=disc[:200],290 anomaly_severity=Severity.CRITICAL,291 )292 result = env.step(action)293 obs = result.observation294 done = result.done295 total_reward += result.reward296 steps.append({"action": f"flag from {doc.document_type}", "reward": result.reward})297 if done:298 return _build_result(task_id, steps, total_reward, result)299 300 # Final verdict301 if not done:302 # Decide verdict based on anomalies found303 has_anomalies = len(obs.flagged_anomalies) > 0304 has_suspicious = any(305 inv.suspicious for inv in obs.investigation_results306 )307 308 if has_anomalies or has_suspicious:309 verdict = VerdictType.FRAUDULENT310 confidence = min(0.95, 0.5 + len(obs.flagged_anomalies) * 0.15)311 else:312 verdict = VerdictType.LEGITIMATE313 confidence = 0.7314 315 action = ITRAction(316 action_type=ActionType.RENDER_VERDICT,317 verdict=verdict,318 confidence=confidence,319 explanation=f"Based on {len(obs.flagged_anomalies)} anomalies detected and {len(obs.investigation_results)} investigations.",320 )321 result = env.step(action)322 obs = result.observation323 done = result.done324 total_reward += result.reward325 steps.append({"action": f"verdict: {verdict.value}", "reward": result.reward})326 327 return _build_result(task_id, steps, total_reward, result)328 329 330def _build_result(task_id, steps, total_reward, result):331 return {332 "task_id": task_id,333 "steps": steps,334 "total_reward": total_reward,335 "final_score": result.info.get("final_score", 0.0),336 "num_steps": len(steps),337 }338 339 340def main():341 parser = argparse.ArgumentParser(description="ITR Fraud Detection Baseline")342 parser.add_argument("--local", action="store_true", help="Use local heuristic agent (no API key needed)")343 parser.add_argument("--task", type=str, default=None, help="Run specific task (easy/medium/hard)")344 args = parser.parse_args()345 346 # Fix Windows encoding347 if sys.stdout.encoding and sys.stdout.encoding.lower() != "utf-8":348 sys.stdout.reconfigure(encoding="utf-8", errors="replace") # type: ignore349 350 env = ITREnvironment()351 tasks = [args.task] if args.task else ["easy", "medium", "hard"]352 353 print("=" * 70)354 print(" ITR FRAUD DETECTION - BASELINE INFERENCE")355 print("=" * 70)356 357 all_scores = {}358 359 for task_id in tasks:360 print(f"[START] {task_id}")361 362 if args.local:363 result = run_heuristic_agent(env, task_id)364 else:365 result = run_openai_agent(env, task_id)366 367 score = result["final_score"]368 all_scores[task_id] = score369 370 print(f"[END] {task_id} score={score:.4f}")371 372 # Print step-by-step373 print(f"\n Step-by-step:")374 for i, step in enumerate(result["steps"]):375 action_str = step.get("action", str(step.get("action", "")))376 print(f" {i+1}. {action_str} (reward: {step['reward']:+.4f})")377 378 print(f"\n{'=' * 70}")379 print(" SUMMARY")380 print(f"{'=' * 70}")381 for task_id, score in all_scores.items():382 bar = "#" * int(score * 30) + "." * (30 - int(score * 30))383 print(f" {task_id:8s} [{bar}] {score:.4f}")384 avg = sum(all_scores.values()) / len(all_scores) if all_scores else 0385 print(f"\n Average Score: {avg:.4f}")386 print(f"{'=' * 70}")387 388 389if __name__ == "__main__":390 main()391 