giangbai / chamthi.py
hoangthiencm's picture
Update chamthi.py
bc37b30 verified
Raw
History Blame Contribute Delete
18.1 kB
import cv2
import numpy as np
import imutils
from imutils.contours import sort_contours
import json
import base64
import logging
import os
import uuid
import google.generativeai as genai
from fastapi import APIRouter, File, UploadFile, Form
from supabase import create_client, Client
# Cấu hình Log
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/omr", tags=["ChamThiOMR"])
# --- CẤU HÌNH GEMINI API ---
GEMINI_API_KEY = os.environ.get("GEMINI_API_KEY")
if GEMINI_API_KEY:
genai.configure(api_key=GEMINI_API_KEY)
# --- KẾT NỐI SUPABASE ---
SUPABASE_URL = os.environ.get("SUPABASE_URL")
SUPABASE_KEY = os.environ.get("SUPABASE_KEY")
supabase: Client = None
if SUPABASE_URL and SUPABASE_KEY:
try:
supabase = create_client(SUPABASE_URL, SUPABASE_KEY)
logger.info("✅ OMR Module: Đã kết nối Supabase")
except Exception as e:
logger.error(f"❌ OMR Module: Lỗi kết nối Supabase: {e}")
# --- BIẾN TOÀN CỤC (CACHE RAM) ---
sessions = {}
# --- CÁC HÀM XỬ LÝ ẢNH ---
def order_points(pts):
"""Sắp xếp 4 điểm góc"""
rect = np.zeros((4, 2), dtype="float32")
s = pts.sum(axis=1)
rect[0] = pts[np.argmin(s)]
rect[2] = pts[np.argmax(s)]
diff = np.diff(pts, axis=1)
rect[1] = pts[np.argmin(diff)]
rect[3] = pts[np.argmax(diff)]
return rect
def four_point_transform(image, pts):
"""Warp ảnh"""
rect = order_points(pts)
(tl, tr, br, bl) = rect
maxWidth = max(int(np.linalg.norm(br - bl)), int(np.linalg.norm(tr - tl)))
maxHeight = max(int(np.linalg.norm(tr - br)), int(np.linalg.norm(tl - bl)))
dst = np.array([[0, 0],[maxWidth - 1, 0],[maxWidth - 1, maxHeight - 1],[0, maxHeight - 1]], dtype="float32")
M = cv2.getPerspectiveTransform(rect, dst)
return cv2.warpPerspective(image, M, (maxWidth, maxHeight))
def read_bubbles(roi, cols, rows, draw_on_me=None, offset=(0,0), bubble_thresh=55):
"""Hàm đọc bong bóng. bubble_thresh=55: Ngưỡng thấp để bắt được bút chì mờ/mực xanh."""
gray = cv2.cvtColor(roi, cv2.COLOR_BGR2GRAY)
thresh = cv2.threshold(gray, 0, 255, cv2.THRESH_BINARY_INV | cv2.THRESH_OTSU)[1]
cnts = cv2.findContours(thresh.copy(), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cnts = imutils.grab_contours(cnts)
bubbles = []
h_img, w_img = roi.shape[:2]
min_w = w_img // (cols * 6)
for c in cnts:
(x, y, w, h) = cv2.boundingRect(c)
ar = w / float(h)
if w >= min_w and 0.5 <= ar <= 1.6:
bubbles.append(c)
if not bubbles: return None
try:
bubbles = sort_contours(bubbles, method="left-to-right")[0]
columns = []
current_col = []
prev_x = -1000
for c in bubbles:
(x, y, w, h) = cv2.boundingRect(c)
if x - prev_x > w:
if current_col:
current_col = sort_contours(current_col, method="top-to-bottom")[0]
columns.append(current_col)
current_col = [c]
prev_x = x + w/2
else:
current_col.append(c)
if current_col:
current_col = sort_contours(current_col, method="top-to-bottom")[0]
columns.append(current_col)
if len(columns) > cols: columns = columns[:cols]
result_str = ""
for col in columns:
filled_idx = -1
max_pixel = 0
for i, c in enumerate(col):
mask = np.zeros(thresh.shape, dtype="uint8")
cv2.drawContours(mask, [c], -1, 255, -1)
mask = cv2.bitwise_and(thresh, thresh, mask=mask)
total = cv2.countNonZero(mask)
if draw_on_me is not None:
(x, y, w, h) = cv2.boundingRect(c)
cv2.rectangle(draw_on_me, (x+offset[0], y+offset[1]), (x+w+offset[0], y+h+offset[1]), (0, 255, 0), 1)
if total > max_pixel:
max_pixel = total
filled_idx = i
if max_pixel > bubble_thresh:
result_str += str(filled_idx) if filled_idx < 10 else "?"
if draw_on_me is not None and filled_idx != -1 and filled_idx < len(col):
c = col[filled_idx]
(x, y, w, h) = cv2.boundingRect(c)
cv2.circle(draw_on_me, (int(x+w/2+offset[0]), int(y+h/2+offset[1])), int(w/2), (0, 0, 255), 3)
else:
result_str += "?"
return result_str
except Exception as e:
logger.error(f"Sort Error: {e}")
return None
# --- AI VISION HELPER ---
def read_handwriting_with_gemini(roi_image):
"""Dùng Gemini Flash để đọc SBD/Mã đề viết tay"""
if not GEMINI_API_KEY: return None, None
try:
_, buffer = cv2.imencode('.jpg', roi_image)
model = genai.GenerativeModel('gemini-1.5-flash')
prompt = """
Analyze this OMR sheet header.
Identify the handwritten Student ID (SBD - usually 6 digits) and Exam Code (Mã đề - usually 3 or 4 digits).
Ignore the printed text or bubbles.
Return JSON format: {"sbd": "...", "code": "..."}
If you cannot find them, return empty strings.
"""
response = model.generate_content([{'mime_type': 'image/jpeg', 'data': buffer.tobytes()}, prompt])
text = response.text.strip()
clean_text = text
if "json" in clean_text:
clean_text = clean_text.replace("```json", "").replace("```", "").strip()
elif "{" in clean_text:
start = clean_text.find("{")
end = clean_text.rfind("}") + 1
if start != -1 and end != -1: clean_text = clean_text[start:end]
data = json.loads(clean_text)
return data.get("sbd", ""), data.get("code", "")
except Exception as e:
logger.error(f"AI Vision Error: {e}")
return None, None
def process_omr(image_bytes, all_keys, use_ai=False, pre_warped=False):
try:
nparr = np.frombuffer(image_bytes, np.uint8)
image = cv2.imdecode(nparr, cv2.IMREAD_COLOR)
if pre_warped:
warped = image
if warped.shape[1] != 1600: warped = imutils.resize(warped, width=1600)
else:
if image.shape[1] > 1600: image = imutils.resize(image, width=1600)
gray = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)
blurred = cv2.GaussianBlur(gray, (5, 5), 0)
edged = cv2.Canny(blurred, 75, 200)
cnts = cv2.findContours(edged.copy(), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cnts = imutils.grab_contours(cnts)
docCnt = None
if len(cnts) > 0:
cnts = sorted(cnts, key=cv2.contourArea, reverse=True)
for c in cnts:
peri = cv2.arcLength(c, True)
approx = cv2.approxPolyDP(c, 0.02 * peri, True)
if len(approx) == 4: docCnt = approx; break
if docCnt is not None:
warped = four_point_transform(image, docCnt.reshape(4, 2))
else:
warped = image
h, w = warped.shape[:2]
draw_img = warped.copy()
# === 1. XỬ LÝ ID (SBD + MÃ ĐỀ) ===
split_x = int(w * 0.36)
id_roi_y = int(h * 0.12)
left_roi = warped[id_roi_y:h-50, 0:split_x]
id_result = read_bubbles(left_roi, 10, 10, draw_on_me=draw_img, offset=(0, id_roi_y), bubble_thresh=55)
sbd = "AI_READ"
code = "DEFAULT"
if id_result and len(id_result) >= 6:
sbd_omr = id_result[:6].replace("?", "")
code_omr = ""
if len(id_result) >= 10: code_omr = id_result[6:10].replace("?", "")
elif len(id_result) >= 9: code_omr = id_result[6:].replace("?", "")
sbd = sbd_omr
code = code_omr if code_omr else "DEFAULT"
# --- AI VISION LAYER ---
if use_ai:
header_roi = warped[0:int(h*0.4), 0:split_x]
ai_sbd, ai_code = read_handwriting_with_gemini(header_roi)
if ai_sbd and (len(sbd) < 6 or "?" in id_result[:6]): sbd = ai_sbd
if ai_code and (code == "DEFAULT" or len(code) < 3): code = ai_code
if not sbd: sbd = "UNKNOWN"
if not code: code = "DEFAULT"
# === 2. XỬ LÝ CÂU HỎI ===
ans_roi_y = int(h * 0.20)
avail_codes = list(all_keys.keys())
if code not in avail_codes and avail_codes:
if "DEFAULT" in avail_codes: code = "DEFAULT"
else: code = str(avail_codes[0])
key_data = all_keys.get(code, {})
total_q = len(key_data)
num_cols_page = 4 if total_q > 30 else 3
col_width = (w - split_x) // num_cols_page
score = 0
correct_count = 0
for c_idx in range(num_cols_page):
c_x_start = split_x + c_idx * col_width
c_roi = warped[ans_roi_y:h-50, c_x_start : c_x_start + col_width]
gray_c = cv2.cvtColor(c_roi, cv2.COLOR_BGR2GRAY)
thresh_c = cv2.threshold(gray_c, 0, 255, cv2.THRESH_BINARY_INV | cv2.THRESH_OTSU)[1]
cnts_c = cv2.findContours(thresh_c, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
cnts_c = imutils.grab_contours(cnts_c)
bubbles_c = []
for b in cnts_c:
(bx, by, bw, bh) = cv2.boundingRect(b)
if bw > 5 and bw < 100: bubbles_c.append(b)
if not bubbles_c: continue
try: bubbles_c = sort_contours(bubbles_c, method="top-to-bottom")[0]
except: continue
rows = []
temp_row = []
prev_y = -1000
for b in bubbles_c:
(bx, by, bw, bh) = cv2.boundingRect(b)
if by - prev_y > bh * 0.8:
if temp_row:
if len(temp_row) == 4:
temp_row = sort_contours(temp_row, method="left-to-right")[0]
rows.append(temp_row)
temp_row = [b]
prev_y = by + bh/2
else: temp_row.append(b)
if len(temp_row) == 4:
temp_row = sort_contours(temp_row, method="left-to-right")[0]
rows.append(temp_row)
for r_idx, row in enumerate(rows):
questions_per_col = 10
q_num = c_idx * questions_per_col + r_idx + 1
if q_num > total_q: break
filled_opt = -1
max_p = 0
for o_idx, b in enumerate(row):
mask = np.zeros(thresh_c.shape, dtype="uint8")
cv2.drawContours(mask, [b], -1, 255, -1)
mask = cv2.bitwise_and(thresh_c, thresh_c, mask=mask)
total = cv2.countNonZero(mask)
(bx, by, bw, bh) = cv2.boundingRect(b)
gx, gy = c_x_start + bx, ans_roi_y + by
cv2.rectangle(draw_img, (gx, gy), (gx+bw, gy+bh), (0, 255, 0), 1)
if total > max_p:
max_p = total
filled_opt = o_idx
correct_char = key_data.get(q_num, key_data.get(str(q_num)))
if correct_char:
correct_idx = ord(correct_char) - 65
if max_p > 55:
b = row[filled_opt]
(bx, by, bw, bh) = cv2.boundingRect(b)
gx, gy = c_x_start + bx, ans_roi_y + by
if filled_opt == correct_idx:
correct_count += 1
cv2.circle(draw_img, (gx+bw//2, gy+bh//2), bw//2, (0, 255, 0), -1)
else:
cv2.circle(draw_img, (gx+bw//2, gy+bh//2), bw//2, (0, 0, 255), -1)
if correct_idx < len(row):
b_correct = row[correct_idx]
(bx_c, by_c, bw_c, bh_c) = cv2.boundingRect(b_correct)
gx_c, gy_c = c_x_start + bx_c, ans_roi_y + by_c
cv2.circle(draw_img, (gx_c+bw_c//2, gy_c+bh_c//2), 5, (255, 0, 0), -1)
if total_q > 0: score = round((correct_count / total_q) * 10, 2)
cv2.rectangle(draw_img, (0, 0), (w, 100), (255, 255, 255), -1)
cv2.putText(draw_img, f"SBD: {sbd} | CODE: {code}", (20, 40), cv2.FONT_HERSHEY_SIMPLEX, 1, (0,0,0), 2)
cv2.putText(draw_img, f"DIEM: {score} ({correct_count}/{total_q})", (20, 90), cv2.FONT_HERSHEY_SIMPLEX, 1.5, (0,0,255), 3)
_, buffer = cv2.imencode('.jpg', draw_img)
b64 = base64.b64encode(buffer).decode('utf-8')
return {
"student_id": sbd, "exam_code": code, "score": score,
"correct_count": correct_count, "total_questions": total_q,
"wrong_count": total_q - correct_count, "image_base64": b64,
"status": "success",
"ai_used": use_ai
}
except Exception as e:
logger.error(f"OMR Error: {str(e)}")
return {"score": 0, "error": str(e)}
# --- DB HELPERS ---
def save_result_to_db(data, filename, session_id):
if not supabase: return
try:
payload = {
"session_id": session_id,
"student_id": data.get("student_id"),
"exam_code": data.get("exam_code"),
"score": data.get("score"),
"correct_count": data.get("correct_count"),
"wrong_count": data.get("wrong_count"),
"total_questions": data.get("total_questions"),
"filename": filename
}
supabase.table("omr_results").insert(payload).execute()
except Exception as e: logger.error(f"DB Save Error: {e}")
def get_key_from_db(session_id):
if not supabase: return {}
try:
res = supabase.table("omr_sessions").select("answer_key_json").eq("session_id", session_id).execute()
if res.data: return json.loads(res.data[0]['answer_key_json'])
except: pass
return {}
# --- API ENDPOINTS ---
@router.post("/create_session")
async def create_session(name: str = Form(...), answer_key_json: str = Form(...)):
session_id = str(uuid.uuid4())[:8].upper()
try: all_keys = json.loads(answer_key_json)
except: all_keys = {}
sessions[session_id] = {"keys": all_keys, "results": []}
if supabase:
try: supabase.table("omr_sessions").insert({"session_id": session_id, "session_name": name, "answer_key_json": answer_key_json}).execute()
except Exception as e: logger.error(f"Create Session DB Error: {e}")
return {"session_id": session_id, "name": name}
@router.post("/grade")
async def grade_exam(file: UploadFile = File(...), answer_key_json: str = Form(...), handwriting: str = Form(None), pre_warped: str = Form(None)):
try: all_keys = json.loads(answer_key_json)
except: return {"error": "JSON Error"}
content = await file.read()
use_ai = (handwriting == "true")
is_pre_warped = (pre_warped == "true")
result = process_omr(content, all_keys, use_ai=use_ai, pre_warped=is_pre_warped)
save_result_to_db(result, file.filename, "Direct_Upload")
return result
@router.post("/grade_mobile")
async def grade_mobile(session_id: str = Form(...), file: UploadFile = File(...), pre_warped: str = Form(None)):
if session_id not in sessions:
db_keys = get_key_from_db(session_id)
if db_keys: sessions[session_id] = {"keys": db_keys, "results": []}
else: return {"error": "Session not found"}
content = await file.read()
is_pre_warped = (pre_warped == "true")
result = process_omr(content, sessions[session_id]["keys"], use_ai=True, pre_warped=is_pre_warped)
result["filename"] = file.filename
sessions[session_id]["results"].append(result)
save_result_to_db(result, file.filename, session_id)
return result
@router.get("/poll_results/{session_id}")
async def poll_results(session_id: str):
if session_id not in sessions: return []
res = sessions[session_id]["results"]
sessions[session_id]["results"] = []
return res
@router.get("/history")
def get_history():
if not supabase: return []
try: return supabase.table("omr_sessions").select("*").order("created_at", desc=True).execute().data
except: return []
@router.get("/history/{session_id}")
def get_history_detail(session_id: str):
if not supabase: return []
try: return supabase.table("omr_results").select("*").eq("session_id", session_id).order("created_at", desc=True).execute().data
except: return []
@router.put("/update_session/{session_id}")
async def update_session(session_id: str, name: str = Form(...)):
if supabase:
try:
supabase.table("omr_sessions").update({"session_name": name}).eq("session_id", session_id).execute()
return {"status": "updated", "name": name}
except Exception as e: return {"error": str(e)}
return {"status": "local_updated"}
@router.delete("/delete_session/{session_id}")
async def delete_session(session_id: str):
if session_id in sessions: del sessions[session_id]
if supabase:
try:
supabase.table("omr_results").delete().eq("session_id", session_id).execute()
supabase.table("omr_sessions").delete().eq("session_id", session_id).execute()
return {"status": "deleted"}
except Exception as e: return {"error": str(e)}
return {"status": "local_deleted"}
@router.post("/init_session")
async def init_session(answer_key_json: str = Form(...)):
return await create_session(name="Quick Session", answer_key_json=answer_key_json)