Download server.py from VishwaK03/Error-Driven-Learning-Reinforcement-Engine: direct link, hf CLI and curl.
- Browser
- Download file 15.1 kB
-
https://huggingface.co/spaces/VishwaK03/Error-Driven-Learning-Reinforcement-Engine/resolve/main/server.py
- Command line
-
hf download hf://spaces/VishwaK03/Error-Driven-Learning-Reinforcement-Engine/server.py
-
curl -L -o server.py https://huggingface.co/spaces/VishwaK03/Error-Driven-Learning-Reinforcement-Engine/resolve/main/server.py
15.1 kB
| from fastapi import FastAPI | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from pydantic import BaseModel | |
| from transformers import AutoTokenizer, AutoModelForSequenceClassification | |
| import torch | |
| import torch.nn.functional as F | |
| from pymongo import MongoClient | |
| import os | |
| import re | |
| import json | |
| import hashlib | |
| import google.generativeai as genai | |
| from dotenv import load_dotenv | |
| from datetime import datetime | |
| # --- CONFIGURATION & SECURITY --- | |
| # Load secrets from .env file | |
| load_dotenv() | |
| MONGO_URI = os.getenv("MONGO_URI") | |
| GENAI_API_KEY = os.getenv("GENAI_API_KEY") | |
| if not GENAI_API_KEY or not MONGO_URI: | |
| print("β ERROR: Missing GENAI_API_KEY or MONGO_URI in .env file!") | |
| # Configure Gemini | |
| genai.configure(api_key=GENAI_API_KEY) | |
| # Change this path if your model is located elsewhere | |
| MODEL_PATH = "./edlre_final_model" | |
| app = FastAPI() | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], # Allows all origins (React port 5173) | |
| allow_credentials=True, | |
| allow_methods=["*"], # Allows all methods (POST, GET, OPTIONS, etc.) | |
| allow_headers=["*"], # Allows all headers | |
| ) | |
| # --- GLOBAL VARIABLES --- | |
| model = None | |
| tokenizer = None | |
| db_connected = False | |
| logs_collection = None | |
| generated_tutorials_collection = None | |
| users_collection = None | |
| intervention_db = {} | |
| labels = {0: 'Syntax Error', 1: 'Semantic Error', 2: 'Logical Error', 3: 'No Error'} | |
| # --- π§ GENAI FUNCTION (Now with Supervisor Validation) --- | |
| def ask_llm_for_intervention(code_snippet, error_class): | |
| print(f"π€ Connecting to Gemini for dynamic {error_class} help...") | |
| try: | |
| valid_model = None | |
| for m in genai.list_models(): | |
| if 'generateContent' in m.supported_generation_methods: | |
| valid_model = m.name | |
| if 'flash' in m.name or 'pro' in m.name: | |
| break | |
| if not valid_model: | |
| raise Exception("No valid models found for this API key.") | |
| print(f"π€ Using model: {valid_model}") | |
| model_gen = genai.GenerativeModel(valid_model) | |
| # π NEW SUPERVISOR PROMPT: Gemini can now disagree with CodeBERT! | |
| prompt = f""" | |
| You are an expert C Tutor supervisor. A smaller AI model flagged this code as a '{error_class}'. | |
| CODE: | |
| {code_snippet} | |
| TASK 1: Verify if the code ACTUALLY has a C programming error. | |
| TASK 2: If the code is perfectly valid C code, you MUST set "total_errors" to 0. | |
| TASK 3: If it DOES have an error, explain it using the JSON structure. | |
| Return ONLY a raw JSON object. Do not use markdown blocks. | |
| JSON Structure: | |
| {{ | |
| "level_1": "π‘ Hint: A short, vague hint. (Watch video)", | |
| "level_2": "β οΈ Error: Explain the bug specifically. (Watch video)", | |
| "level_3": "π Fix: Tell them exactly how to fix it. (Watch video)", | |
| "title": "Short Descriptive Title", | |
| "concept": "Explain the underlying C concept.", | |
| "fix": "Direct fix instruction.", | |
| "bad": "The problematic snippet", | |
| "good": "The corrected snippet", | |
| "error_line": 5, | |
| "total_errors": 1 | |
| }} | |
| Note: If the code is correct, set total_errors to 0. | |
| """ | |
| response = model_gen.generate_content(prompt) | |
| text = response.text | |
| # Clean text | |
| clean_text = re.sub(r'```json|```', '', text).strip() | |
| json_match = re.search(r'\{.*\}', clean_text, re.DOTALL) | |
| if json_match: | |
| data = json.loads(json_match.group(0)) | |
| print(f" β GenAI Success! (Found {data.get('total_errors', 1)} errors)") | |
| return data | |
| except Exception as e: | |
| print(f" β GenAI Error: {e}") | |
| return { | |
| "level_1": f"π‘ Hint: Check your {error_class} logic. (Watch video)", | |
| "level_2": f"β οΈ Error: The NeuroMentor AI detected a {error_class}. (Watch video)", | |
| "level_3": "π Fix: Review your syntax and logic. (Watch video)", | |
| "title": f"C {error_class}", | |
| "concept": f"A {error_class} happens when the code doesn't match the required logic or rules.", | |
| "fix": "Review the relevant sections of your C code.", | |
| "bad": code_snippet[:50] + "...", | |
| "good": f"// Refer to C documentation for {error_class}", | |
| "error_line": -1, | |
| "total_errors": 1 | |
| } | |
| async def startup_event(): | |
| # ADDED users_collection to global list here: | |
| global model, tokenizer, db_connected, logs_collection, generated_tutorials_collection, users_collection, intervention_db | |
| print("π SERVER STARTING...") | |
| try: | |
| if os.path.exists("interventions.json"): | |
| with open("interventions.json", "r", encoding="utf-8") as f: | |
| intervention_db = json.load(f) | |
| print("1οΈβ£ Interventions: β SUCCESS!") | |
| else: | |
| intervention_db = {} | |
| except Exception as e: | |
| intervention_db = {} | |
| try: | |
| client = MongoClient(MONGO_URI, serverSelectionTimeoutMS=2000) | |
| client.admin.command('ping') | |
| db = client["neuromentor_db"] | |
| logs_collection = db["intervention_logs"] | |
| generated_tutorials_collection = db["generated_tutorials"] | |
| users_collection = db["users"] # <--- ADDED THIS NEW COLLECTION | |
| db_connected = True | |
| print("2οΈβ£ MongoDB: β SUCCESS!") | |
| except Exception as e: | |
| db_connected = False | |
| print("2οΈβ£ MongoDB: β CONNECTION FAILED") | |
| try: | |
| tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH) | |
| full_model = AutoModelForSequenceClassification.from_pretrained(MODEL_PATH) | |
| model = torch.quantization.quantize_dynamic(full_model, {torch.nn.Linear}, dtype=torch.qint8) | |
| print("3οΈβ£ AI Model: β SUCCESS!") | |
| except Exception as e: | |
| model = None | |
| print("3οΈβ£ AI Model: β FAILED!") | |
| class CodeRequest(BaseModel): | |
| code: str | |
| user_id: str = "novice_001" | |
| cognitive_state: str = "neutral" | |
| def analyze_code_snapshot(code_snippet): | |
| lines = code_snippet.split('\n') | |
| for i, line in enumerate(lines): | |
| line_num = i + 1 | |
| if re.search(r'(?<!\.)\b\d+\s*/\s*\d+\b(?!\.)', line): return "Logical Error", "integer_division_trap", line_num | |
| if re.search(r'if\s*\(\s*\w+\s*=(?!=)\s*[\w\d]+\s*\)', line): return "Logical Error", "assignment_in_condition", line_num | |
| if re.search(r'==\s*"', line): return "Logical Error", "string_equality_error", line_num | |
| if re.search(r'if\s*\(.*\)\s*;', line): return "Logical Error", "if_semicolon_trap", line_num | |
| if re.search(r'scanf\s*\(\s*"%d"\s*,\s*[a-zA-Z0-9_]+\s*\)', line): return "Syntax Error", "scanf_missing_ampersand", line_num | |
| if "malloc" in line and "free" not in code_snippet: return "Semantic Error", "memory_leak", line_num | |
| if "NULL" in line and re.search(r'\*\w+\s*=', line): return "Semantic Error", "null_pointer_dereference", line_num | |
| return None, None, -1 | |
| async def predict_error(request: CodeRequest): | |
| code = request.code | |
| user_id = request.user_id | |
| state = request.cognitive_state.lower() # π Grab the state from VS Code! | |
| # 1. Rules | |
| label, tag, error_line = analyze_code_snapshot(code) | |
| source = "Syllabus Rule Engine" | |
| confidence = 100.0 | |
| total_errors = 1 if label else 0 | |
| # 2. AI Model (CodeBERT) | |
| if not label: | |
| if model: | |
| try: | |
| inputs = tokenizer(code, return_tensors="pt", truncation=True, max_length=512) | |
| with torch.no_grad(): outputs = model(**inputs) | |
| probs = F.softmax(outputs.logits, dim=-1) | |
| conf, pred_class = torch.max(probs, dim=-1) | |
| label = labels[pred_class.item()] | |
| confidence = conf.item() * 100 | |
| source = "NeuroMentor AI" | |
| tag = "general_error" if label != "No Error" else "correct_code" | |
| except: label, tag = "No Error", "correct_code" | |
| else: label, tag = "No Error", "correct_code" | |
| full_data = None | |
| # --------------------------------------------------------- | |
| # π§ 3. THE COGNITIVE MAPPING ENGINE π§ | |
| # Override scaffolding level based on real-time brain state! | |
| # --------------------------------------------------------- | |
| if "confused" in state: | |
| level = 1 # Level 1: Maximum help, full explanations | |
| elif "neutral" in state or "relaxed" in state: | |
| level = 2 # Level 2: Standard hint | |
| elif "focused" in state or "active_thinking" in state: | |
| level = 3 # Level 3: Minimal nudge to keep them in the flow! | |
| else: | |
| level = 2 # Default fallback | |
| print(f"π§ State: {state} -> Assigned Scaffolding Level: {level}") | |
| level_key = f"level_{level}" | |
| # 4. Fetch / Generate Intervention | |
| if tag in intervention_db and tag != "general_error": | |
| full_data = intervention_db[tag] | |
| full_data["error_line"] = error_line | |
| full_data["total_errors"] = total_errors | |
| elif label != "No Error": | |
| full_data = ask_llm_for_intervention(code, label) | |
| if full_data: | |
| error_line = full_data.get("error_line", -1) | |
| total_errors = full_data.get("total_errors", 1) | |
| # --- π‘οΈ AI SELF-CORRECTION LAYER π‘οΈ --- | |
| if total_errors == 0: | |
| print(" π‘οΈ AI Supervisor Override: Code is actually correct!") | |
| label = "No Error" | |
| tag = "correct_code" | |
| full_data = None # This triggers the "Great Job" UI | |
| source = "NeuroMentor AI Supervisor" | |
| confidence = 100.0 | |
| else: | |
| source = "NeuroMentor AI " | |
| if db_connected: | |
| try: | |
| generated_tutorials_collection.insert_one({ | |
| "error_class": label, | |
| "original_code": code, | |
| "generated_tutorial": full_data, | |
| "timestamp": datetime.now() | |
| }) | |
| except Exception as e: pass | |
| # Only use fallback if it's ACTUALLY an error | |
| if not full_data and label != "No Error": | |
| full_data = { | |
| "level_1": "Hint: Check logic.", "level_2": "Error detected.", "level_3": "Fix syntax.", | |
| "title": "Unknown Error", "concept": "Check logic.", "fix": "Debug.", "bad": "", "good": "", | |
| "error_line": error_line, "total_errors": total_errors | |
| } | |
| recommendation = full_data.get(level_key, full_data.get("level_1")) if full_data else "" | |
| # 5. Log everything to MongoDB | |
| if db_connected: | |
| try: | |
| logs_collection.insert_one({ | |
| "user_id": user_id, | |
| "cognitive_state": state, # π Logging the exact state! | |
| "code": code[:100], | |
| "error": label, | |
| "tag": tag, | |
| "source": source, | |
| "level": level, # π Logging the dynamically calculated level! | |
| "error_line": error_line, | |
| "tutorial": full_data, | |
| "timestamp": datetime.now() | |
| }) | |
| except: pass | |
| print(f"π {label} | Tag: {tag} | Line: {error_line} | Total: {total_errors}") | |
| # 6. Return response to VS Code | |
| return { | |
| "error_type": label, | |
| "tag": tag, | |
| "recommendation": recommendation, | |
| "tutorial": full_data if full_data else None, | |
| "source": source, | |
| "confidence": f"{confidence:.2f}%", | |
| "error_line": error_line, | |
| "total_errors": total_errors | |
| } | |
| # --- π AUTHENTICATION & DASHBOARD API π --- | |
| class UserAuth(BaseModel): | |
| username: str | |
| password: str | |
| def hash_password(password: str): | |
| return hashlib.sha256(password.encode()).hexdigest() | |
| async def signup(user: UserAuth): | |
| if not db_connected: return {"error": "Database offline"} | |
| existing_user = users_collection.find_one({"username": user.username}) | |
| if existing_user: return {"error": "Username already exists"} | |
| new_user = { | |
| "username": user.username, | |
| "password": hash_password(user.password), | |
| "created_at": datetime.now() | |
| } | |
| users_collection.insert_one(new_user) | |
| return {"success": True, "message": "Account created successfully!"} | |
| async def login(user: UserAuth): | |
| if not db_connected: return {"error": "Database offline"} | |
| db_user = users_collection.find_one({"username": user.username}) | |
| if not db_user or db_user["password"] != hash_password(user.password): | |
| return {"error": "Invalid username or password"} | |
| return {"success": True, "username": user.username} | |
| async def get_dashboard(user_id: str): | |
| if not db_connected: return {"error": "Database offline"} | |
| # 1. Total Files Analyzed | |
| total_files = logs_collection.count_documents({"user_id": user_id}) | |
| # 2. Most Frequent Error | |
| pipeline = [ | |
| {"$match": {"user_id": user_id, "error": {"$ne": "No Error"}}}, | |
| {"$group": {"_id": "$error", "count": {"$sum": 1}}}, | |
| {"$sort": {"count": -1}}, | |
| {"$limit": 1} | |
| ] | |
| frequent_error_cursor = list(logs_collection.aggregate(pipeline)) | |
| most_frequent = frequent_error_cursor[0]["_id"] if frequent_error_cursor else "None yet" | |
| # 3. Recent Logs | |
| recent_cursor = logs_collection.find({"user_id": user_id, "error": {"$ne": "No Error"}}).sort("timestamp", -1).limit(10) | |
| recent_logs = [] | |
| for log in recent_cursor: | |
| # Gracefully handle timezone differences just in case | |
| try: | |
| time_diff = datetime.now() - log["timestamp"] | |
| except TypeError: | |
| time_diff = datetime.now(timezone.utc) - log["timestamp"] | |
| minutes_ago = int(time_diff.total_seconds() / 60) | |
| time_str = f"{minutes_ago} mins ago" if minutes_ago < 60 else f"{int(minutes_ago/60)} hours ago" | |
| recent_logs.append({ | |
| "id": str(log["_id"]), | |
| "error": log["error"], | |
| "tag": log["tag"], | |
| "level": log.get("level", 1), | |
| "cognitive_state": log.get("cognitive_state", "neutral"), | |
| "time": time_str, | |
| "tutorial": log.get("tutorial", None) | |
| }) | |
| # 4. π‘οΈ SAFELY define current_state (Bulletproof fix!) | |
| current_state = "Tracking..." | |
| if len(recent_logs) > 0: | |
| current_state = recent_logs[0]["cognitive_state"] | |
| return { | |
| "totalFiles": total_files, | |
| "mostFrequentError": most_frequent, | |
| "cognitiveState": current_state, | |
| "recentLogs": recent_logs | |
| } | |
| if __name__ == "__main__": | |
| import uvicorn | |
| uvicorn.run(app, host="127.0.0.1", port=8080) |