Phoenix1410 commited on
Commit
2dd6a59
·
verified ·
1 Parent(s): 13e0403

Hopefully the final one.

Browse files
Files changed (1) hide show
  1. main.py +187 -237
main.py CHANGED
@@ -1,17 +1,29 @@
1
  from __future__ import annotations
 
 
 
 
 
 
 
2
  from fastapi import FastAPI, UploadFile, File, Form, HTTPException, Security, Depends, status
3
  from fastapi.middleware.cors import CORSMiddleware
4
  from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
5
  from fastapi.responses import StreamingResponse
6
  from pydantic import BaseModel
7
- from transformers import pipeline, AutoTokenizer, AutoModelForSequenceClassification, T5Tokenizer, T5ForConditionalGeneration
 
 
 
 
 
 
8
  from sentence_transformers import SentenceTransformer, CrossEncoder
9
  from sklearn.metrics.pairwise import cosine_similarity
10
  from langchain_groq import ChatGroq
11
  from langchain_core.prompts import ChatPromptTemplate
12
- import os
13
- import fitz
14
- import jwt
15
  import numpy as np
16
  import torch
17
  import asyncio
@@ -20,21 +32,28 @@ import json
20
  from dotenv import load_dotenv
21
  import sys
22
  import io
 
23
 
24
- os.environ["USE_TF"] = "0"
25
- os.environ["USE_TORCH"] = "1"
26
- # Ensure UTF-8 stdout on Windows
27
  if sys.stdout and hasattr(sys.stdout, 'buffer'):
28
  sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding='utf-8', errors='replace')
29
  if sys.stderr and hasattr(sys.stderr, 'buffer'):
30
  sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding='utf-8', errors='replace')
31
 
32
- # 1. Setup App & Configuration
33
  load_dotenv()
34
  app = FastAPI(title="LexGuard AI Core")
35
 
36
- from database import db
37
- from models import UserSync, Timeline, TimelineEvent, Discrepancy, ComparativeAnalysisResult
 
 
 
 
 
 
 
 
38
 
39
  # --- CORS CONFIGURATION ---
40
  app.add_middleware(
@@ -49,7 +68,7 @@ app.add_middleware(
49
  MODEL_DIR = "./lexguard_model"
50
  ABS_MODEL_PATH = os.path.abspath(MODEL_DIR)
51
  GROQ_API_KEY = os.getenv("GROQ_API_KEY")
52
- CLERK_JWKS_URL = os.getenv("CLERK_JWKS_URL") # Set in your .env for production signature verification
53
 
54
  # --- SECURITY (CLERK AUTH) ---
55
  security = HTTPBearer()
@@ -61,7 +80,6 @@ def verify_clerk_token(credentials: HTTPAuthorizationCredentials = Security(secu
61
  return "demo_user"
62
  try:
63
  if CLERK_JWKS_URL:
64
- # Requires 'cryptography' package installed
65
  jwks_client = jwt.PyJWKClient(CLERK_JWKS_URL)
66
  signing_key = jwks_client.get_signing_key_from_jwt(token)
67
  payload = jwt.decode(
@@ -71,34 +89,29 @@ def verify_clerk_token(credentials: HTTPAuthorizationCredentials = Security(secu
71
  options={"verify_exp": True}
72
  )
73
  else:
74
- # Fallback/Prototype mode (No signature check)
75
  payload = jwt.decode(token, options={"verify_signature": False})
76
 
77
  user_id = payload.get("sub")
78
- if not user_id:
79
- return "authenticated_user"
80
- return user_id
81
  except Exception as e:
82
- print(f"[AUTH WARNING] Clerk token decode notice: {e}")
83
- # In prototype mode without JWKS, gracefully treat valid format attempts as authenticated
84
  return "clerk_user"
85
 
86
  # --- LOAD MODELS ---
87
-
88
- # Dynamic device mapping (CUDA vs CPU fallback)
89
  device_id = 0 if torch.cuda.is_available() else -1
90
  device_name = "GPU (CUDA)" if device_id == 0 else "CPU"
91
 
92
- # 1. THE SNIPER (DistilRoBERTa - Risk Detection)
93
  print(f"Loading Sniper Model from: {ABS_MODEL_PATH}")
94
  try:
95
- model = AutoModelForSequenceClassification.from_pretrained(ABS_MODEL_PATH, local_files_only=True)
96
- tokenizer = AutoTokenizer.from_pretrained(ABS_MODEL_PATH, local_files_only=True)
97
- sniper = pipeline("text-classification", model=model, tokenizer=tokenizer, device=device_id)
98
  print(f"[OK] Sniper Model Loaded ({device_name})")
99
  except Exception as e:
100
- print(f"[ERROR] Error loading Sniper: {e}")
101
- exit()
 
102
 
103
  # 2. THE SCOUT (Sentence-BERT - Semantic Search)
104
  print("Loading Scout Model (Semantic Search)...")
@@ -107,23 +120,35 @@ try:
107
  scout = SentenceTransformer('all-MiniLM-L6-v2', device=scout_device)
108
  print(f"[OK] Scout Model Loaded ({scout_device.upper()})")
109
  except Exception as e:
110
- print(f"[ERROR] Error loading Scout: {e}")
111
- exit()
112
 
113
- # 3. THE SENTINEL (Local Cross-Encoder NLI for Testimony Contradiction Filtering)
114
  print("Loading Sentinel Model (Local NLI Triage)...")
115
  try:
116
  nli_device = "cuda" if torch.cuda.is_available() else "cpu"
117
- # Lightweight DeBERTa-v3 (~140MB) optimized for Premise-Hypothesis Contradiction classification
118
  nli_classifier = CrossEncoder("cross-encoder/nli-deberta-v3-small", device=nli_device)
119
  print(f"[OK] Sentinel NLI Model Loaded ({nli_device.upper()})")
120
  except Exception as e:
121
- print(f"[WARNING] Sentinel NLI load warning: {e}. Falling back to cosine thresholding.")
122
  nli_classifier = None
123
 
124
- # 4. THE ANALYST (Groq Reasoning Engine)
 
 
 
 
 
 
 
 
 
 
 
 
 
125
  if not GROQ_API_KEY:
126
- print("WARNING: GROQ_API_KEY not found in environment.")
127
 
128
  DEFAULT_MODEL = os.getenv("GROQ_MODEL", "llama-3.3-70b-versatile")
129
  analyst = ChatGroq(
@@ -136,7 +161,7 @@ print(f"[OK] Analyst Model Loaded: {DEFAULT_MODEL}")
136
  # --- HELPER FUNCTIONS ---
137
 
138
  def extract_text_from_pdf(file_bytes: bytes) -> list[str]:
139
- """Parses PDF bytes into structured, bite-sized legal clause chunks."""
140
  doc = fitz.open(stream=file_bytes, filetype="pdf")
141
  text_chunks = []
142
 
@@ -154,19 +179,16 @@ def extract_text_from_pdf(file_bytes: bytes) -> list[str]:
154
  if len(clean_text) > 250:
155
  sub_clauses = re.split(r'(?<=[.!?])\s+(?=[A-Z0-9("])|(?<=;)\s+(?=[A-Z0-9("])', clean_text)
156
  current_chunk = ""
157
-
158
  for clause in sub_clauses:
159
  clause_str = clause.strip()
160
  if not clause_str:
161
  continue
162
-
163
  if len(current_chunk) + len(clause_str) < 250:
164
  current_chunk = f"{current_chunk} {clause_str}".strip() if current_chunk else clause_str
165
  else:
166
  if current_chunk and len(current_chunk) > 40:
167
  text_chunks.append(current_chunk)
168
  current_chunk = clause_str
169
-
170
  if current_chunk and len(current_chunk) > 40:
171
  text_chunks.append(current_chunk)
172
  else:
@@ -174,72 +196,15 @@ def extract_text_from_pdf(file_bytes: bytes) -> list[str]:
174
 
175
  return text_chunks
176
 
177
- def segment_transcript_statements(transcript: str) -> list[str]:
178
- """
179
- Splits deposition/testimony text into standalone semantic units.
180
- Preserves Q&A turns where present.
181
- """
182
- qa_blocks = re.findall(r'(Q:.*?)(?=(?:Q:|$))', transcript, flags=re.DOTALL | re.IGNORECASE)
183
- if qa_blocks:
184
- return [re.sub(r'\s+', ' ', b).strip() for b in qa_blocks if len(b.strip()) > 15]
185
-
186
- # Fallback: split by sentence boundaries
187
- sentences = re.split(r'(?<=[.!?])\s+', transcript)
188
- return [s.strip() for s in sentences if len(s.strip()) > 20]
189
-
190
- def local_triage_testimony_pairs(client_statements: list[str], accused_statements: list[str], top_k: int = 2) -> list[dict]:
191
- """
192
- Local Vector + NLI pass:
193
- 1. Encodes statements with Scout (SentenceTransformer).
194
- 2. Finds semantically aligned statement pairs.
195
- 3. Runs Sentinel (CrossEncoder NLI) to filter for high-probability contradictions.
196
- """
197
- if not client_statements or not accused_statements:
198
- return []
199
-
200
- client_vecs = scout.encode(client_statements, convert_to_numpy=True, normalize_embeddings=True)
201
- accused_vecs = scout.encode(accused_statements, convert_to_numpy=True, normalize_embeddings=True)
202
-
203
- sim_matrix = np.dot(client_vecs, accused_vecs.T)
204
- candidate_pairs = []
205
- cross_encoder_pairs = []
206
-
207
- for c_idx, c_stmt in enumerate(client_statements):
208
- top_accused_indices = np.argsort(sim_matrix[c_idx])[::-1][:top_k]
209
- for a_idx in top_accused_indices:
210
- similarity = float(sim_matrix[c_idx][a_idx])
211
- # Focus on topically correlated claims
212
- if similarity >= 0.25:
213
- a_stmt = accused_statements[a_idx]
214
- candidate_pairs.append({
215
- "client_claim": c_stmt,
216
- "accused_claim": a_stmt,
217
- "similarity": similarity
218
- })
219
- cross_encoder_pairs.append((c_stmt, a_stmt))
220
-
221
- if not candidate_pairs:
222
- return []
223
-
224
- # If CrossEncoder is loaded, score contradiction probability locally
225
- if nli_classifier and cross_encoder_pairs:
226
- # CrossEncoder classes: 0: Contradiction, 1: Entailment, 2: Neutral
227
- nli_scores = nli_classifier.predict(cross_encoder_pairs, apply_softmax=True)
228
- flagged = []
229
- for pair_dict, score_dist in zip(candidate_pairs, nli_scores):
230
- contra_score = float(score_dist[0])
231
- pair_dict["contradiction_prob"] = contra_score
232
- # Keep items with plausible contradiction or divergence
233
- if contra_score > 0.35:
234
- flagged.append(pair_dict)
235
- flagged.sort(key=lambda x: x["contradiction_prob"], reverse=True)
236
- return flagged[:6]
237
-
238
- # Fallback to top similarity pairs if NLI is inactive
239
- return candidate_pairs[:6]
240
-
241
- async def process_analyst_evaluation(clause: str, user_rule: str | None, source_str: str, index: int, pred_score: float, risk_type: str):
242
- """Asynchronous wrapper to query the LLM concurrently for a flagged clause."""
243
  system_msg = f"""You are an elite legal auditor.
244
  Detection Reason: {source_str}
245
  User's Constraint Rule: {user_rule if user_rule else "None"}
@@ -255,7 +220,8 @@ Task:
255
  ])
256
 
257
  try:
258
- ai_response = await (prompt | analyst).ainvoke({
 
259
  "system_msg": system_msg,
260
  "clause_text": clause
261
  })
@@ -273,13 +239,85 @@ Task:
273
  "source": source_str
274
  }
275
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
276
  # --- API ENDPOINTS ---
277
 
278
  @app.get("/")
279
  def health_check():
280
  return {
281
  "status": "LexGuard Brain is Online 🧠",
282
- "models": ["Sniper", "Scout", "Sentinel (NLI)", "Analyst"]
283
  }
284
 
285
  @app.post("/users/sync")
@@ -297,8 +335,8 @@ async def sync_user(user_data: UserSync):
297
  )
298
  return {"status": "User synced successfully", "clerk_id": user_data.clerk_id}
299
  except Exception as e:
300
- print(f"Error syncing user: {e}")
301
- raise HTTPException(status_code=500, detail=f"Database sync failed: {str(e)}")
302
 
303
  @app.post("/analyze_document")
304
  async def analyze_document(
@@ -306,9 +344,9 @@ async def analyze_document(
306
  user_rule: str = Form(None),
307
  user_id: str = Depends(verify_clerk_token)
308
  ):
 
309
  print(f"[UPLOADING] User {user_id} uploading: {file.filename}")
310
 
311
- # A. Parse PDF
312
  try:
313
  content = await file.read()
314
  clauses = extract_text_from_pdf(content)
@@ -323,37 +361,27 @@ async def analyze_document(
323
  "results": []
324
  }
325
 
326
- print(f"[SCANNING] Scanning {len(clauses)} clauses...")
327
-
328
- # B. THE SNIPER PASS (Batch Classification)
329
  label_map = {"LABEL_0": "Safe", "LABEL_1": "Termination", "LABEL_2": "Non-Compete"}
330
 
331
- for idx, clause in enumerate(clauses):
332
- tokens = tokenizer.encode(clause, add_special_tokens=True)
333
- if len(tokens) > 512:
334
- print(f"[WARNING] Clause {idx} exceeds 512 tokens. Text will be truncated for Sniper.")
335
-
336
- sniper_preds = sniper(clauses, batch_size=8, truncation=True)
337
 
338
- # C. THE SCOUT PASS (Semantic Search)
339
  semantic_matches = set()
340
-
341
- if user_rule and len(user_rule.strip()) > 5:
342
- print(f"[SCOUT] Scout searching for rule: '{user_rule}'")
343
  rule_vec = scout.encode([user_rule])
344
  clause_vecs = scout.encode(clauses)
345
-
346
  sim_scores = cosine_similarity(rule_vec, clause_vecs)[0]
347
- top_indices = np.argsort(sim_scores)[-3:]
348
-
349
  for idx in top_indices:
350
  if sim_scores[idx] > 0.30:
351
  semantic_matches.add(int(idx))
352
- print(f" -> Match at Clause {idx} (Score: {sim_scores[idx]:.2f})")
353
 
354
- # D. CONCURRENT AGGREGATION & GENAI
355
  analysis_tasks = []
356
-
357
  for i, (clause, pred) in enumerate(zip(clauses, sniper_preds)):
358
  label_str = pred['label']
359
  risk_type = label_map.get(label_str, "Safe")
@@ -380,10 +408,7 @@ async def analyze_document(
380
  )
381
  analysis_tasks.append(task)
382
 
383
- if analysis_tasks:
384
- results = await asyncio.gather(*analysis_tasks)
385
- else:
386
- results = []
387
 
388
  return {
389
  "filename": file.filename,
@@ -393,93 +418,21 @@ async def analyze_document(
393
  "results": results
394
  }
395
 
396
- # --- TESTIMONY VALIDATOR (Hybrid Local NLI + Few-Shot Groq Map-Reduce) ---
397
-
398
- print("Loading Interrogator Model (Local T5 QG)...")
399
- try:
400
- from transformers import T5Tokenizer, T5ForConditionalGeneration
401
- qg_tokenizer = T5Tokenizer.from_pretrained("doc2query/msmarco-t5-base-v1", local_files_only=False)
402
- qg_model = T5ForConditionalGeneration.from_pretrained("doc2query/msmarco-t5-base-v1", local_files_only=False)
403
- qg_model.to(device_name.lower() if device_name == "cuda" else "cpu")
404
- print(f"[OK] Interrogator QG Model Loaded")
405
- except Exception as e:
406
- print(f"[ERROR] Interrogator QG failed to load: {e}")
407
-
408
- # --- HELPER: LOCAL QUESTION GENERATION ---
409
- def generate_local_questions(text: str, max_questions: int = 5) -> list[str]:
410
- """Uses local T5 to generate relevant interrogation questions from the text."""
411
- # Grab the most substantial sentences to generate questions
412
- sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', text) if len(s.strip()) > 40]
413
- target_sentences = sentences[:max_questions] # Keep it token-cheap
414
-
415
- questions = set()
416
- for sentence in target_sentences:
417
- input_ids = qg_tokenizer.encode(sentence, return_tensors='pt').to(qg_model.device)
418
- outputs = qg_model.generate(
419
- input_ids=input_ids,
420
- max_length=64,
421
- do_sample=True,
422
- top_k=10,
423
- num_return_sequences=1
424
- )
425
- q = qg_tokenizer.decode(outputs[0], skip_special_tokens=True)
426
- if "?" in q:
427
- # Clean formatting
428
- q = q.split("?")[0] + "?"
429
- questions.add(q.capitalize())
430
-
431
- return list(questions)
432
-
433
- # --- HELPER: GROQ PARALLEL QUERYING ---
434
- async def extract_answers_from_text(questions: list[str], text: str, role: str) -> dict:
435
- """Asks Groq to answer the list of questions strictly based on the provided text."""
436
-
437
- # We ask Groq to return a strict JSON dictionary { "Question 1": "Answer", ... }
438
- system_prompt = f"""You are a neutral fact-extractor.
439
- Read the {role} testimony carefully.
440
- Answer the provided questions strictly using the facts stated in the text.
441
- If the text does not contain the answer, your exact response MUST be: "Not mentioned."
442
- Output valid JSON where keys are the exact questions, and values are the answers."""
443
-
444
- prompt = ChatPromptTemplate.from_messages([
445
- ("system", system_prompt),
446
- ("human", "QUESTIONS:\n{questions}\n\nTEXT:\n{text}")
447
- ])
448
-
449
- try:
450
- response = await analyst.ainvoke({
451
- "questions": json.dumps(questions),
452
- "text": text
453
- })
454
-
455
- # Parse JSON output from Groq
456
- raw_content = response.content
457
- # Strip markdown json blocks if present
458
- if raw_content.startswith("```json"):
459
- raw_content = raw_content[7:-3]
460
-
461
- return json.loads(raw_content.strip())
462
- except Exception as e:
463
- print(f"[ERROR] Groq extraction failed for {role}: {e}")
464
- return {q: "Error retrieving data" for q in questions}
465
-
466
- # --- ENDPOINT: STREAMING COMPARATOR ---
467
  @app.post("/stream_compare_testimonies")
468
  async def stream_compare_testimonies(
469
  client_file: UploadFile = File(None),
470
  client_text: str = Form(None),
471
  accused_file: UploadFile = File(None),
472
  accused_text: str = Form(None),
473
- # user_id: str = Depends(verify_clerk_token) # Uncomment in prod
474
  ):
475
  """
476
- Streaming Endpoint (Server-Sent Events).
477
- 1. Local T5 generates questions from Client text.
478
- 2. Groq queries BOTH texts in parallel.
479
- 3. Streams results row-by-row to the frontend for the flashing UI.
480
  """
481
 
482
- # 1. Resolve inputs
483
  async def resolve_input(file: UploadFile, text: str) -> str:
484
  if file and file.filename:
485
  content = await file.read()
@@ -491,42 +444,42 @@ async def stream_compare_testimonies(
491
  c_content = await resolve_input(client_file, client_text)
492
  a_content = await resolve_input(accused_file, accused_text)
493
 
494
- if not c_content or not a_content:
495
- raise HTTPException(status_code=400, detail="Missing testimony data.")
496
 
497
- # 2. Generator for Server-Sent Events (SSE)
498
  async def event_generator():
499
  try:
500
- # Event: Booting up / Generating Questions
501
- yield f"data: {json.dumps({'status': 'initializing', 'msg': 'Local AI extracting core interrogations...'})}\n\n"
502
- await asyncio.sleep(0.5)
503
 
504
- # Local T5 Execution
505
  questions = generate_local_questions(c_content, max_questions=6)
506
- if not questions:
507
- questions = ["What are the main events described?", "Who was involved?"]
508
-
509
  yield f"data: {json.dumps({'status': 'questions_ready', 'msg': f'Generated {len(questions)} factual queries.'})}\n\n"
 
 
 
 
510
 
511
- # Event: Querying Groq in Parallel
512
- yield f"data: {json.dumps({'status': 'querying', 'msg': 'Querying independent testimonies via Groq LPU...'})}\n\n"
513
-
514
- # Fire both Groq prompts SIMULTANEOUSLY
515
  client_answers, accused_answers = await asyncio.gather(
516
- extract_answers_from_text(questions, c_content, "Client"),
517
- extract_answers_from_text(questions, a_content, "Accused")
518
  )
519
 
520
- # Event: Streaming the Matrix Comparison Row-by-Row
521
  for q in questions:
522
  ans_c = client_answers.get(q, "Not mentioned.")
523
  ans_a = accused_answers.get(q, "Not mentioned.")
524
 
525
- # Unbiased logic check
526
- is_omission = "Not mentioned" in ans_c or "Not mentioned" in ans_a
527
- match_status = "Incomplete Event" if is_omission else "Event details do not align"
528
- if ans_c.lower() == ans_a.lower() and not is_omission:
529
  match_status = "Accounts Align"
 
 
 
 
530
 
531
  payload = {
532
  "status": "flashing_pair",
@@ -536,18 +489,15 @@ async def stream_compare_testimonies(
536
  "match_status": match_status
537
  }
538
 
539
- # Stream the payload to Vercel instantly
540
  yield f"data: {json.dumps(payload)}\n\n"
541
-
542
- # Pause for 2 seconds to allow the Frontend UI to flash it on screen
543
  await asyncio.sleep(2.0)
544
 
545
- # Close Stream
546
- yield f"data: {json.dumps({'status': 'done', 'msg': 'Forensic evaluation complete.'})}\n\n"
547
 
548
  except Exception as e:
 
549
  yield f"data: {json.dumps({'status': 'error', 'msg': str(e)})}\n\n"
550
 
551
- # Return HTTP Stream
552
- return StreamingResponse(event_generator(), media_type="text/event-stream")
553
-
 
1
  from __future__ import annotations
2
+ import os
3
+
4
+ # 1. Force PyTorch mode: Prevents Transformers from scanning TensorFlow/Keras 3
5
+ os.environ["USE_TF"] = "0"
6
+ os.environ["USE_TORCH"] = "1"
7
+ os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"
8
+
9
  from fastapi import FastAPI, UploadFile, File, Form, HTTPException, Security, Depends, status
10
  from fastapi.middleware.cors import CORSMiddleware
11
  from fastapi.security import HTTPBearer, HTTPAuthorizationCredentials
12
  from fastapi.responses import StreamingResponse
13
  from pydantic import BaseModel
14
+ from transformers import (
15
+ pipeline,
16
+ AutoTokenizer,
17
+ AutoModelForSequenceClassification,
18
+ T5Tokenizer,
19
+ T5ForConditionalGeneration
20
+ )
21
  from sentence_transformers import SentenceTransformer, CrossEncoder
22
  from sklearn.metrics.pairwise import cosine_similarity
23
  from langchain_groq import ChatGroq
24
  from langchain_core.prompts import ChatPromptTemplate
25
+ import fitz # PyMuPDF
26
+ import jwt # PyJWT
 
27
  import numpy as np
28
  import torch
29
  import asyncio
 
32
  from dotenv import load_dotenv
33
  import sys
34
  import io
35
+ from typing import Optional, List, Dict, Any
36
 
37
+ # Ensure UTF-8 stdout on Windows / Container environments
 
 
38
  if sys.stdout and hasattr(sys.stdout, 'buffer'):
39
  sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding='utf-8', errors='replace')
40
  if sys.stderr and hasattr(sys.stderr, 'buffer'):
41
  sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding='utf-8', errors='replace')
42
 
43
+ # 2. Setup App & Configuration
44
  load_dotenv()
45
  app = FastAPI(title="LexGuard AI Core")
46
 
47
+ # Safe database & model import with fallback
48
+ try:
49
+ from database import db
50
+ from models import UserSync, Timeline, TimelineEvent, Discrepancy, ComparativeAnalysisResult
51
+ except ImportError:
52
+ class UserSync(BaseModel):
53
+ clerk_id: str
54
+ email: Optional[str] = None
55
+ name: Optional[str] = None
56
+ created_at: Optional[str] = None
57
 
58
  # --- CORS CONFIGURATION ---
59
  app.add_middleware(
 
68
  MODEL_DIR = "./lexguard_model"
69
  ABS_MODEL_PATH = os.path.abspath(MODEL_DIR)
70
  GROQ_API_KEY = os.getenv("GROQ_API_KEY")
71
+ CLERK_JWKS_URL = os.getenv("CLERK_JWKS_URL")
72
 
73
  # --- SECURITY (CLERK AUTH) ---
74
  security = HTTPBearer()
 
80
  return "demo_user"
81
  try:
82
  if CLERK_JWKS_URL:
 
83
  jwks_client = jwt.PyJWKClient(CLERK_JWKS_URL)
84
  signing_key = jwks_client.get_signing_key_from_jwt(token)
85
  payload = jwt.decode(
 
89
  options={"verify_exp": True}
90
  )
91
  else:
 
92
  payload = jwt.decode(token, options={"verify_signature": False})
93
 
94
  user_id = payload.get("sub")
95
+ return user_id if user_id else "authenticated_user"
 
 
96
  except Exception as e:
97
+ print(f"[AUTH WARNING] Clerk token notice: {e}")
 
98
  return "clerk_user"
99
 
100
  # --- LOAD MODELS ---
 
 
101
  device_id = 0 if torch.cuda.is_available() else -1
102
  device_name = "GPU (CUDA)" if device_id == 0 else "CPU"
103
 
104
+ # 1. THE SNIPER (DistilRoBERTa - Contract Risk Detection)
105
  print(f"Loading Sniper Model from: {ABS_MODEL_PATH}")
106
  try:
107
+ sniper_model = AutoModelForSequenceClassification.from_pretrained(ABS_MODEL_PATH, local_files_only=True)
108
+ sniper_tokenizer = AutoTokenizer.from_pretrained(ABS_MODEL_PATH, local_files_only=True)
109
+ sniper = pipeline("text-classification", model=sniper_model, tokenizer=sniper_tokenizer, device=device_id)
110
  print(f"[OK] Sniper Model Loaded ({device_name})")
111
  except Exception as e:
112
+ print(f"[WARNING] Sniper Model load warning: {e}")
113
+ sniper = None
114
+ sniper_tokenizer = None
115
 
116
  # 2. THE SCOUT (Sentence-BERT - Semantic Search)
117
  print("Loading Scout Model (Semantic Search)...")
 
120
  scout = SentenceTransformer('all-MiniLM-L6-v2', device=scout_device)
121
  print(f"[OK] Scout Model Loaded ({scout_device.upper()})")
122
  except Exception as e:
123
+ print(f"[WARNING] Scout Model load warning: {e}")
124
+ scout = None
125
 
126
+ # 3. THE SENTINEL (Cross-Encoder NLI)
127
  print("Loading Sentinel Model (Local NLI Triage)...")
128
  try:
129
  nli_device = "cuda" if torch.cuda.is_available() else "cpu"
 
130
  nli_classifier = CrossEncoder("cross-encoder/nli-deberta-v3-small", device=nli_device)
131
  print(f"[OK] Sentinel NLI Model Loaded ({nli_device.upper()})")
132
  except Exception as e:
133
+ print(f"[WARNING] Sentinel NLI load warning: {e}")
134
  nli_classifier = None
135
 
136
+ # 4. THE INTERROGATOR (Local T5 Question Generator)
137
+ print("Loading Interrogator Model (Local T5 QG)...")
138
+ try:
139
+ qg_device = "cuda" if torch.cuda.is_available() else "cpu"
140
+ qg_tokenizer = T5Tokenizer.from_pretrained("doc2query/msmarco-t5-base-v1", local_files_only=False)
141
+ qg_model = T5ForConditionalGeneration.from_pretrained("doc2query/msmarco-t5-base-v1", local_files_only=False)
142
+ qg_model.to(qg_device)
143
+ print(f"[OK] Interrogator QG Model Loaded ({qg_device.upper()})")
144
+ except Exception as e:
145
+ print(f"[WARNING] Interrogator QG Model load warning: {e}")
146
+ qg_tokenizer = None
147
+ qg_model = None
148
+
149
+ # 5. THE ANALYST (Groq Reasoning Engine)
150
  if not GROQ_API_KEY:
151
+ print("[WARNING] GROQ_API_KEY not found in environment.")
152
 
153
  DEFAULT_MODEL = os.getenv("GROQ_MODEL", "llama-3.3-70b-versatile")
154
  analyst = ChatGroq(
 
161
  # --- HELPER FUNCTIONS ---
162
 
163
  def extract_text_from_pdf(file_bytes: bytes) -> list[str]:
164
+ """Parses PDF bytes into structured legal clause chunks."""
165
  doc = fitz.open(stream=file_bytes, filetype="pdf")
166
  text_chunks = []
167
 
 
179
  if len(clean_text) > 250:
180
  sub_clauses = re.split(r'(?<=[.!?])\s+(?=[A-Z0-9("])|(?<=;)\s+(?=[A-Z0-9("])', clean_text)
181
  current_chunk = ""
 
182
  for clause in sub_clauses:
183
  clause_str = clause.strip()
184
  if not clause_str:
185
  continue
 
186
  if len(current_chunk) + len(clause_str) < 250:
187
  current_chunk = f"{current_chunk} {clause_str}".strip() if current_chunk else clause_str
188
  else:
189
  if current_chunk and len(current_chunk) > 40:
190
  text_chunks.append(current_chunk)
191
  current_chunk = clause_str
 
192
  if current_chunk and len(current_chunk) > 40:
193
  text_chunks.append(current_chunk)
194
  else:
 
196
 
197
  return text_chunks
198
 
199
+ async def process_analyst_evaluation(
200
+ clause: str,
201
+ user_rule: Optional[str],
202
+ source_str: str,
203
+ index: int,
204
+ pred_score: float,
205
+ risk_type: str
206
+ ):
207
+ """Evaluates single document contract risks asynchronously."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
208
  system_msg = f"""You are an elite legal auditor.
209
  Detection Reason: {source_str}
210
  User's Constraint Rule: {user_rule if user_rule else "None"}
 
220
  ])
221
 
222
  try:
223
+ chain = prompt | analyst
224
+ ai_response = await chain.ainvoke({
225
  "system_msg": system_msg,
226
  "clause_text": clause
227
  })
 
239
  "source": source_str
240
  }
241
 
242
+ def generate_local_questions(text: str, max_questions: int = 6) -> list[str]:
243
+ """Generates sharp factual interrogation queries using local T5."""
244
+ if not qg_model or not qg_tokenizer:
245
+ return [
246
+ "What specific date and time are mentioned?",
247
+ "Who were the key individuals present?",
248
+ "What agreements or formal commitments were made?",
249
+ "Where did the described meeting or event take place?"
250
+ ]
251
+
252
+ sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', text) if len(s.strip()) > 35]
253
+ target_sentences = sentences[:max_questions]
254
+ questions = set()
255
+
256
+ for sentence in target_sentences:
257
+ input_ids = qg_tokenizer.encode(sentence, return_tensors='pt').to(qg_model.device)
258
+ outputs = qg_model.generate(
259
+ input_ids=input_ids,
260
+ max_length=64,
261
+ do_sample=True,
262
+ top_k=10,
263
+ num_return_sequences=1
264
+ )
265
+ q = qg_tokenizer.decode(outputs[0], skip_special_tokens=True)
266
+ if "?" in q:
267
+ clean_q = q.split("?")[0].strip() + "?"
268
+ questions.add(clean_q.capitalize())
269
+
270
+ if not questions:
271
+ return [
272
+ "What specific date and time are mentioned?",
273
+ "Who were the key individuals present?",
274
+ "What agreements or formal commitments were made?"
275
+ ]
276
+
277
+ return list(questions)
278
+
279
+ async def extract_answers_from_text(questions: list[str], text: str, role: str) -> dict:
280
+ """Asks Groq to answer questions strictly using facts from the text."""
281
+ system_prompt = f"""You are a neutral forensic fact-extractor.
282
+ Read the {role} testimony carefully.
283
+ Answer each of the provided questions strictly using the explicit facts stated in the text.
284
+ If the text does not contain the answer, your exact response MUST be: "Not mentioned."
285
+ You MUST respond ONLY with a valid JSON object where keys are the exact questions and values are the answers.
286
+ Do not include conversational preamble, explanations, or markdown code blocks."""
287
+
288
+ prompt = ChatPromptTemplate.from_messages([
289
+ ("system", system_prompt),
290
+ ("human", "QUESTIONS TO ANSWER:\n{questions_str}\n\nTESTIMONY TEXT:\n{text}")
291
+ ])
292
+
293
+ chain = prompt | analyst
294
+
295
+ try:
296
+ questions_formatted = "\n".join([f"- {q}" for q in questions])
297
+ response = await chain.ainvoke({
298
+ "questions_str": questions_formatted,
299
+ "text": text
300
+ })
301
+
302
+ raw_content = response.content.strip()
303
+
304
+ # Bulletproof JSON boundary parsing (prevents markdown or text wrappers from crashing loads)
305
+ json_match = re.search(r"\{.*\}", raw_content, re.DOTALL)
306
+ if json_match:
307
+ return json.loads(json_match.group(0))
308
+ return json.loads(raw_content)
309
+
310
+ except Exception as e:
311
+ print(f"[ERROR] Groq extraction failed for {role}: {e}")
312
+ return {q: "Not mentioned." for q in questions}
313
+
314
  # --- API ENDPOINTS ---
315
 
316
  @app.get("/")
317
  def health_check():
318
  return {
319
  "status": "LexGuard Brain is Online 🧠",
320
+ "models": ["Sniper", "Scout", "Sentinel (NLI)", "Interrogator (T5)", "Analyst"]
321
  }
322
 
323
  @app.post("/users/sync")
 
335
  )
336
  return {"status": "User synced successfully", "clerk_id": user_data.clerk_id}
337
  except Exception as e:
338
+ print(f"[ERROR] Syncing user failed: {e}")
339
+ return {"status": "User sync skipped (standalone mode)"}
340
 
341
  @app.post("/analyze_document")
342
  async def analyze_document(
 
344
  user_rule: str = Form(None),
345
  user_id: str = Depends(verify_clerk_token)
346
  ):
347
+ """Scans contracts using Sniper (DistilRoBERTa), Scout (Sentence-BERT), and Analyst."""
348
  print(f"[UPLOADING] User {user_id} uploading: {file.filename}")
349
 
 
350
  try:
351
  content = await file.read()
352
  clauses = extract_text_from_pdf(content)
 
361
  "results": []
362
  }
363
 
 
 
 
364
  label_map = {"LABEL_0": "Safe", "LABEL_1": "Termination", "LABEL_2": "Non-Compete"}
365
 
366
+ # Run Sniper Classification if model available
367
+ if sniper and sniper_tokenizer:
368
+ sniper_preds = sniper(clauses, batch_size=8, truncation=True)
369
+ else:
370
+ sniper_preds = [{'label': 'LABEL_0', 'score': 1.0} for _ in clauses]
 
371
 
372
+ # Run Scout Search
373
  semantic_matches = set()
374
+ if scout and user_rule and len(user_rule.strip()) > 5:
 
 
375
  rule_vec = scout.encode([user_rule])
376
  clause_vecs = scout.encode(clauses)
 
377
  sim_scores = cosine_similarity(rule_vec, clause_vecs)[0]
378
+ top_indices = np.argsort(sim_scores)[-3:]
 
379
  for idx in top_indices:
380
  if sim_scores[idx] > 0.30:
381
  semantic_matches.add(int(idx))
 
382
 
383
+ # Aggregate and analyze risks with Groq concurrently
384
  analysis_tasks = []
 
385
  for i, (clause, pred) in enumerate(zip(clauses, sniper_preds)):
386
  label_str = pred['label']
387
  risk_type = label_map.get(label_str, "Safe")
 
408
  )
409
  analysis_tasks.append(task)
410
 
411
+ results = await asyncio.gather(*analysis_tasks) if analysis_tasks else []
 
 
 
412
 
413
  return {
414
  "filename": file.filename,
 
418
  "results": results
419
  }
420
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
421
  @app.post("/stream_compare_testimonies")
422
  async def stream_compare_testimonies(
423
  client_file: UploadFile = File(None),
424
  client_text: str = Form(None),
425
  accused_file: UploadFile = File(None),
426
  accused_text: str = Form(None),
427
+ user_id: str = Depends(verify_clerk_token)
428
  ):
429
  """
430
+ Symmetric Streaming Testimony Validator (SSE).
431
+ 1. Local T5 extracts factual interrogation questions from reference text.
432
+ 2. Groq queries BOTH testimonies concurrently via parallel execution.
433
+ 3. Yields real-time comparison events to power the dynamic flashing UI.
434
  """
435
 
 
436
  async def resolve_input(file: UploadFile, text: str) -> str:
437
  if file and file.filename:
438
  content = await file.read()
 
444
  c_content = await resolve_input(client_file, client_text)
445
  a_content = await resolve_input(accused_file, accused_text)
446
 
447
+ if not c_content.strip() or not a_content.strip():
448
+ raise HTTPException(status_code=400, detail="Missing testimony inputs for Party A or Party B.")
449
 
 
450
  async def event_generator():
451
  try:
452
+ # 1. Initialization Event
453
+ yield f"data: {json.dumps({'status': 'initializing', 'msg': 'Extracting core factual interrogations locally...'})}\n\n"
454
+ await asyncio.sleep(0.3)
455
 
456
+ # 2. Local T5 Question Generation
457
  questions = generate_local_questions(c_content, max_questions=6)
 
 
 
458
  yield f"data: {json.dumps({'status': 'questions_ready', 'msg': f'Generated {len(questions)} factual queries.'})}\n\n"
459
+ await asyncio.sleep(0.3)
460
+
461
+ # 3. Parallel Querying via Groq LPU
462
+ yield f"data: {json.dumps({'status': 'querying', 'msg': 'Querying parallel accounts simultaneously...'})}\n\n"
463
 
 
 
 
 
464
  client_answers, accused_answers = await asyncio.gather(
465
+ extract_answers_from_text(questions, c_content, "Party A"),
466
+ extract_answers_from_text(questions, a_content, "Party B")
467
  )
468
 
469
+ # 4. Stream Matrix Results Row-by-Row
470
  for q in questions:
471
  ans_c = client_answers.get(q, "Not mentioned.")
472
  ans_a = accused_answers.get(q, "Not mentioned.")
473
 
474
+ # Neutral evaluation criteria
475
+ is_omission = "not mentioned" in ans_c.lower() or "not mentioned" in ans_a.lower()
476
+
477
+ if ans_c.strip().lower() == ans_a.strip().lower() and not is_omission:
478
  match_status = "Accounts Align"
479
+ elif is_omission:
480
+ match_status = "Incomplete Event"
481
+ else:
482
+ match_status = "Event details do not align"
483
 
484
  payload = {
485
  "status": "flashing_pair",
 
489
  "match_status": match_status
490
  }
491
 
 
492
  yield f"data: {json.dumps(payload)}\n\n"
493
+ # Hold window for the frontend to animate and display
 
494
  await asyncio.sleep(2.0)
495
 
496
+ # 5. Complete Stream
497
+ yield f"data: {json.dumps({'status': 'done', 'msg': 'Forensic cross-examination complete.'})}\n\n"
498
 
499
  except Exception as e:
500
+ print(f"[STREAM ERROR] {e}")
501
  yield f"data: {json.dumps({'status': 'error', 'msg': str(e)})}\n\n"
502
 
503
+ return StreamingResponse(event_generator(), media_type="text/event-stream")