Phoenix1410 commited on
Commit
5dee3a6
Β·
verified Β·
1 Parent(s): d819566

Update main.py

Browse files
Files changed (1) hide show
  1. main.py +34 -32
main.py CHANGED
@@ -6,7 +6,6 @@ from langchain_groq import ChatGroq
6
  from langchain_core.prompts import ChatPromptTemplate
7
  import os
8
  import torch
9
- import jwt
10
  import fitz # PyMuPDF
11
  from dotenv import load_dotenv
12
  from typing import Dict, Optional
@@ -19,7 +18,7 @@ from database import db
19
  load_dotenv()
20
  app = FastAPI(title="LexGuard AI Core")
21
 
22
- # --- CORS CONFIGURATION (CRITICAL FOR VERCEL & REACT) ---
23
  origins = [
24
  "http://localhost:3000",
25
  "http://127.0.0.1:3000",
@@ -36,33 +35,43 @@ app.add_middleware(
36
  allow_headers=["*"],
37
  )
38
 
39
- # --- MODEL PATHS ---
40
  MODEL_DIR = "./lexguard_model"
41
  ABS_MODEL_PATH = os.path.abspath(MODEL_DIR)
42
- GROQ_API_KEY = os.getenv("GROQ_API_KEY")
 
 
 
 
 
 
 
 
43
 
44
  # 2. Load Models (The Sniper & The Analyst)
45
  print(f"Loading Predictive Model from: {ABS_MODEL_PATH}")
46
  try:
47
- # Force local loading to prevent Hugging Face URL errors
48
  model = AutoModelForSequenceClassification.from_pretrained(
49
  ABS_MODEL_PATH, local_files_only=True
50
  )
51
  tokenizer = AutoTokenizer.from_pretrained(
52
  ABS_MODEL_PATH, local_files_only=True
53
  )
54
- # Automatically switch to CPU (-1) if running on a free-tier non-GPU Hugging Face Space
55
  device = 0 if torch.cuda.is_available() else -1
56
  classifier = pipeline("text-classification", model=model, tokenizer=tokenizer, device=device)
57
  print(f"βœ… Sniper Model Loaded successfully on device: {'GPU' if device == 0 else 'CPU'}")
58
  except Exception as e:
59
- print(f"❌ Error loading model: {e}")
60
- exit()
61
 
 
62
  llm = ChatGroq(
63
  temperature=0,
64
  model_name="llama-3.3-70b-versatile",
65
- groq_api_key=GROQ_API_KEY
 
 
66
  )
67
 
68
  # 3. Helper: PDF to Text Chunks
@@ -79,7 +88,7 @@ def extract_text_from_pdf(file_bytes):
79
  paragraphs = text.split('\n\n')
80
  for p in paragraphs:
81
  clean_p = " ".join(p.split()).strip()
82
- # Filter out tiny page numbers or headers
83
  if len(clean_p) > 50:
84
  text_chunks.append(clean_p)
85
  return text_chunks
@@ -90,7 +99,8 @@ def extract_text_from_pdf(file_bytes):
90
  def health_check():
91
  return {
92
  "status": "LexGuard API is Ready",
93
- "gpu_active": torch.cuda.is_available()
 
94
  }
95
 
96
  @app.post("/users/sync")
@@ -100,7 +110,6 @@ async def sync_user(user: UserSync):
100
  Syncs user from Clerk Frontend to MongoDB Backend.
101
  """
102
  user_data = user.model_dump()
103
- # Upsert based on clerk_id
104
  result = await db.users.update_one(
105
  {"clerk_id": user.clerk_id},
106
  {"$set": user_data},
@@ -112,7 +121,7 @@ async def sync_user(user: UserSync):
112
  @app.post("/analyze_document/")
113
  async def analyze_document(
114
  file: UploadFile = File(...),
115
- user_rule: str = Form(None) # React sends this as FormData
116
  ):
117
  """
118
  Main Endpoint: Receives a PDF file + Optional User Rule.
@@ -140,10 +149,7 @@ async def analyze_document(
140
  # B. The Sniper Pass (Batch Prediction)
141
  print(f"πŸ” Scanning {len(clauses)} clauses...")
142
 
143
- # Map model output to human labels
144
  label_map = {"LABEL_0": "Safe", "LABEL_1": "Termination", "LABEL_2": "Non-Compete"}
145
-
146
- # Run inference
147
  predictions = classifier(clauses, batch_size=8, truncation=True)
148
 
149
  # C. The Filter & Logic Pass
@@ -152,7 +158,6 @@ async def analyze_document(
152
  score = pred['score']
153
  risk_type = label_map.get(label_str, "Safe")
154
 
155
- # LOGIC: We only keep it if it's RISKY OR if User has a specific rule to check
156
  is_risky = risk_type != "Safe"
157
  has_rule = user_rule is not None and len(user_rule.strip()) > 5
158
 
@@ -160,20 +165,17 @@ async def analyze_document(
160
  explanation = "Standard clause."
161
 
162
  # D. The Analyst Pass (GenAI)
163
- if is_risky or (has_rule and i < 20): # Limit rule checking to first 20 clauses for speed
164
-
165
- system_msg = f"""
166
- You are a legal auditor.
167
- Detected Risk Category: {risk_type} (Confidence: {score:.2f})
168
- User's Constraint Rule: {user_rule if user_rule else "None"}
169
-
170
- Task:
171
- 1. Summarize what this clause says in plain English.
172
- 2. If a User Rule exists, explicitly state if this clause violates it.
173
- 3. If it is risky, suggest a 1-sentence edit to make it safer.
174
- """
175
 
176
- # Use parameter placeholder to prevent syntax errors on raw contract brackets
177
  prompt = ChatPromptTemplate.from_messages([
178
  ("system", system_msg),
179
  ("human", "{clause_text}")
@@ -183,10 +185,10 @@ async def analyze_document(
183
  ai_response = (prompt | llm).invoke({"clause_text": clause})
184
  explanation = ai_response.content
185
  except Exception as e:
186
- print(f"Groq API call error: {e}")
 
187
  explanation = "AI Analysis unavailable."
188
 
189
- # Append to results
190
  results.append({
191
  "id": i,
192
  "text": clause,
 
6
  from langchain_core.prompts import ChatPromptTemplate
7
  import os
8
  import torch
 
9
  import fitz # PyMuPDF
10
  from dotenv import load_dotenv
11
  from typing import Dict, Optional
 
18
  load_dotenv()
19
  app = FastAPI(title="LexGuard AI Core")
20
 
21
+ # --- CORS CONFIGURATION ---
22
  origins = [
23
  "http://localhost:3000",
24
  "http://127.0.0.1:3000",
 
35
  allow_headers=["*"],
36
  )
37
 
38
+ # --- MODEL PATHS & API KEYS ---
39
  MODEL_DIR = "./lexguard_model"
40
  ABS_MODEL_PATH = os.path.abspath(MODEL_DIR)
41
+
42
+ # Retrieve and sanitize the Groq API key
43
+ raw_groq_key = os.getenv("GROQ_API_KEY")
44
+ GROQ_API_KEY = raw_groq_key.strip() if raw_groq_key else None
45
+
46
+ if not GROQ_API_KEY:
47
+ print("⚠️ WARNING: GROQ_API_KEY is not detected in environment variables!")
48
+ else:
49
+ print(f"πŸ”‘ GROQ_API_KEY detected (starts with: {GROQ_API_KEY[:8]}...)")
50
 
51
  # 2. Load Models (The Sniper & The Analyst)
52
  print(f"Loading Predictive Model from: {ABS_MODEL_PATH}")
53
  try:
 
54
  model = AutoModelForSequenceClassification.from_pretrained(
55
  ABS_MODEL_PATH, local_files_only=True
56
  )
57
  tokenizer = AutoTokenizer.from_pretrained(
58
  ABS_MODEL_PATH, local_files_only=True
59
  )
60
+ # Dynamically select GPU (0) if CUDA is available, otherwise CPU (-1)
61
  device = 0 if torch.cuda.is_available() else -1
62
  classifier = pipeline("text-classification", model=model, tokenizer=tokenizer, device=device)
63
  print(f"βœ… Sniper Model Loaded successfully on device: {'GPU' if device == 0 else 'CPU'}")
64
  except Exception as e:
65
+ print(f"❌ Error loading predictive model: {e}")
66
+ exit(1)
67
 
68
+ # Initialize Groq LLM client
69
  llm = ChatGroq(
70
  temperature=0,
71
  model_name="llama-3.3-70b-versatile",
72
+ groq_api_key=GROQ_API_KEY,
73
+ request_timeout=60,
74
+ max_retries=2
75
  )
76
 
77
  # 3. Helper: PDF to Text Chunks
 
88
  paragraphs = text.split('\n\n')
89
  for p in paragraphs:
90
  clean_p = " ".join(p.split()).strip()
91
+ # Filter out tiny page numbers or short headers
92
  if len(clean_p) > 50:
93
  text_chunks.append(clean_p)
94
  return text_chunks
 
99
  def health_check():
100
  return {
101
  "status": "LexGuard API is Ready",
102
+ "gpu_active": torch.cuda.is_available(),
103
+ "groq_configured": GROQ_API_KEY is not None
104
  }
105
 
106
  @app.post("/users/sync")
 
110
  Syncs user from Clerk Frontend to MongoDB Backend.
111
  """
112
  user_data = user.model_dump()
 
113
  result = await db.users.update_one(
114
  {"clerk_id": user.clerk_id},
115
  {"$set": user_data},
 
121
  @app.post("/analyze_document/")
122
  async def analyze_document(
123
  file: UploadFile = File(...),
124
+ user_rule: str = Form(None)
125
  ):
126
  """
127
  Main Endpoint: Receives a PDF file + Optional User Rule.
 
149
  # B. The Sniper Pass (Batch Prediction)
150
  print(f"πŸ” Scanning {len(clauses)} clauses...")
151
 
 
152
  label_map = {"LABEL_0": "Safe", "LABEL_1": "Termination", "LABEL_2": "Non-Compete"}
 
 
153
  predictions = classifier(clauses, batch_size=8, truncation=True)
154
 
155
  # C. The Filter & Logic Pass
 
158
  score = pred['score']
159
  risk_type = label_map.get(label_str, "Safe")
160
 
 
161
  is_risky = risk_type != "Safe"
162
  has_rule = user_rule is not None and len(user_rule.strip()) > 5
163
 
 
165
  explanation = "Standard clause."
166
 
167
  # D. The Analyst Pass (GenAI)
168
+ if is_risky or (has_rule and i < 20):
169
+ system_msg = f"""You are a legal auditor.
170
+ Detected Risk Category: {risk_type} (Confidence: {score:.2f})
171
+ User's Constraint Rule: {user_rule if user_rule else "None"}
172
+
173
+ Task:
174
+ 1. Summarize what this clause says in plain English.
175
+ 2. If a User Rule exists, explicitly state if this clause violates it.
176
+ 3. If it is risky, suggest a 1-sentence edit to make it safer."""
 
 
 
177
 
178
+ # Parameterized human prompt prevents template syntax errors on contract brackets
179
  prompt = ChatPromptTemplate.from_messages([
180
  ("system", system_msg),
181
  ("human", "{clause_text}")
 
185
  ai_response = (prompt | llm).invoke({"clause_text": clause})
186
  explanation = ai_response.content
187
  except Exception as e:
188
+ cause = getattr(e, '__cause__', None)
189
+ print(f"❌ Groq API call error: {e} | Underlying cause: {cause}")
190
  explanation = "AI Analysis unavailable."
191
 
 
192
  results.append({
193
  "id": i,
194
  "text": clause,