Team Ai
Apppublic

code-shrish/itr-detection

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py391 linesDownload Raw Back to root
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