JAA-ATS-Tool / src /model_pool.py
saitejatirunagari's picture
perf+feat: parallel resume generation, PDF output, full JD fetching, API timeouts
b15fd58
Raw
History Blame
18.7 kB
"""
ModelPool — 2-phase parallel AI worker pool.
Phase 1 (FAST): Score all jobs quickly → ask only for score (1-10)
Output: ~15 tokens per job → very fast
Batch size: 10 jobs per call
Phase 2 (DETAILED): Full assessment on top-scoring jobs only
Output: ~250 tokens per job
Batch size: 3 jobs per call
4 model workers run simultaneously, grabbing batches from a shared queue.
"""
import json
import re
import time
import queue
import threading
from openai import OpenAI
from colorama import Fore, Style
QUICK_BATCH_SIZE = 10 # Phase 1 LLM batch (unused — replaced by keyword scoring)
DETAILED_BATCH_SIZE = 8 # Phase 2: 8 jobs per LLM call (more = faster, less API calls)
PHASE2_THRESHOLD = 5 # Only do detailed for jobs scoring >= this (keyword threshold)
PHASE2_MAX = 200 # No meaningful cap — process all jobs
class ModelWorker:
"""A single AI model — acts as an independent worker."""
def __init__(self, name: str, model: str, api_key: str, base_url: str, extra_body: dict = None):
self.name = name
self.model = model
self.extra_body = extra_body or {}
# Hard timeout so a hung API call fails fast instead of stalling the pool
self.client = OpenAI(base_url=base_url, api_key=api_key,
timeout=120.0, max_retries=0)
self.jobs_done = 0
self.calls_done = 0
self.errors = 0
def _call(self, messages: list, max_tokens: int, retries: int = 2) -> str:
kwargs = dict(
model=self.model,
messages=messages,
temperature=0.1,
top_p=0.9,
max_tokens=max_tokens,
stream=False,
)
if self.extra_body:
kwargs["extra_body"] = self.extra_body
for attempt in range(retries):
try:
resp = self.client.chat.completions.create(**kwargs)
return resp.choices[0].message.content or ""
except Exception as e:
err_str = str(e)
if "504" in err_str or "502" in err_str or "503" in err_str:
raise RuntimeError(f"Server error ({self.name}): {err_str[:60]}")
if "429" in err_str:
# Rate limit — wait and retry
wait = 15 * (attempt + 1)
time.sleep(wait)
if attempt == retries - 1:
raise RuntimeError(f"Rate limited ({self.name})")
continue
if attempt < retries - 1:
time.sleep(2 ** attempt)
else:
raise
# ── PHASE 1: Quick score (minimal output) ─────────────────────────────────
def quick_score_batch(self, jobs_batch: list[dict], compact_profile: str) -> list[dict]:
"""
Ask only for a score (1-10) per job.
Uses a compact prompt — models respond with arrays like [7,4,9] or ["[7,4,9]"].
"""
n = len(jobs_batch)
lines = "\n".join(
f"JOB {i}: {j.get('title','')} at {j.get('company','')} ({j.get('location','')})"
for i, j in enumerate(jobs_batch, 1)
)
user_msg = (
f"Candidate: {compact_profile}\n\n"
f"{lines}\n\n"
f"Return ONLY a JSON array of {n} integers (1-10 score per job). "
f"Example: [7,4,9]. No explanation."
)
# Generous tokens so model doesn't truncate mid-response
max_tokens = max(200, 20 * n)
try:
text = self._call(
[
{"role": "system", "content": "You are a PM recruiter. Return ONLY valid JSON array. No markdown."},
{"role": "user", "content": user_msg},
],
max_tokens=max_tokens,
)
scores = self._parse_score_array(text, n)
self.calls_done += 1
self.jobs_done += n
return [{"job_index": i + 1, "score": scores[i]} for i in range(n)]
except Exception:
return [{"job_index": i + 1, "score": 5} for i in range(n)]
# ── PHASE 2: Detailed assessment (top jobs only) ───────────────────────────
def detailed_batch(self, jobs_batch: list[dict], compact_profile: str) -> list[dict]:
"""
Full assessment for top-scoring jobs.
Asks for matching skills, gaps, keywords, recommendation.
"""
jobs_text = ""
for i, job in enumerate(jobs_batch, 1):
desc = (job.get("description") or "")[:300].replace("\n", " ")
jobs_text += (
f"\nJOB {i}: {job.get('title','')} at {job.get('company','')} | {job.get('location','')}\n"
f"DESC: {desc}\n"
)
user = (
f"CANDIDATE: {compact_profile}\n"
f"{jobs_text}\n"
f"For each job return a JSON array:\n"
f'[{{"i":1,"score":<1-10>,"pct":<0-100>,"exp":"Good fit|Under-qualified|Over-qualified",'
f'"match":["s1","s2"],"miss":["s1"],"note":"<15 words>",'
f'"kw":["k1","k2","k3"],"pri":"High|Medium|Low"}}]'
)
# ~250 tokens per job
max_tokens = min(350 * len(jobs_batch), 2000)
try:
text = self._call(
[{"role": "system", "content": "PM recruiter. Return ONLY valid JSON array. No markdown."},
{"role": "user", "content": user}],
max_tokens=max_tokens,
)
result = self._parse_json_list(text, len(jobs_batch))
self.calls_done += 1
self.jobs_done += len(jobs_batch)
return result
except Exception:
return [_neutral_detailed(i + 1) for i in range(len(jobs_batch))]
# ── JSON parsers ──────────────────────────────────────────────────────────
def _parse_score_array(self, text: str, expected: int) -> list[int]:
"""
Robust parser for score arrays — handles any format the models return:
[7,4,9,3,5] | {"scores":[7,4,9]} | 7\n4\n9 | Score 1: 7, Score 2: 4 ...
"""
def clamp(n):
try:
return max(1, min(10, int(float(n))))
except Exception:
return 5
def extract_from_list(lst):
nums = []
for item in lst:
if isinstance(item, (int, float)):
nums.append(clamp(item))
elif isinstance(item, str):
# Kimi wraps: ["[7,6,8]"] — unwrap inner string
try:
inner = json.loads(item)
if isinstance(inner, list):
nums.extend(clamp(x) for x in inner if isinstance(x, (int, float)))
elif isinstance(inner, (int, float)):
nums.append(clamp(inner))
except Exception:
# Try extracting digits from the string
found = re.findall(r'\b(10|[1-9])\b', item)
nums.extend(clamp(n) for n in found)
elif isinstance(item, dict):
v = item.get("score") or item.get("s") or item.get("rating") or 5
nums.append(clamp(v))
return nums
# Try JSON parse
try:
result = json.loads(text.strip())
nums = []
if isinstance(result, list):
nums = extract_from_list(result)
elif isinstance(result, dict):
for v in result.values():
if isinstance(v, list):
nums = extract_from_list(v)
break
if not nums:
nums = [clamp(v) for v in result.values() if isinstance(v, (int, float))]
if nums:
while len(nums) < expected: nums.append(5)
return nums[:expected]
except Exception:
pass
# Regex fallback: all standalone integers 1-10
all_nums = re.findall(r'\b(10|[1-9])\b', text)
if all_nums:
nums = [clamp(n) for n in all_nums]
while len(nums) < expected: nums.append(5)
return nums[:expected]
return [5] * expected
def _parse_json_list(self, text: str, expected: int) -> list[dict]:
"""Parse a JSON array from model response."""
# Try direct parse
try:
result = json.loads(text.strip())
if isinstance(result, list):
return result if result else [_neutral_detailed(i + 1) for i in range(expected)]
except Exception:
pass
# Try extracting from code block
m = re.search(r"```(?:json)?\s*([\s\S]+?)```", text)
if m:
try:
result = json.loads(m.group(1))
if isinstance(result, list):
return result
except Exception:
pass
# Extract first [...] block with objects
idx = text.find("[{")
if idx >= 0:
depth = 0
for i, ch in enumerate(text[idx:], idx):
if ch == "[":
depth += 1
elif ch == "]":
depth -= 1
if depth == 0:
try:
result = json.loads(text[idx:i+1])
if isinstance(result, list):
return result
except Exception:
break
return [_neutral_detailed(i + 1) for i in range(expected)]
# ── Helpers ───────────────────────────────────────────────────────────────────
def _neutral_detailed(idx: int) -> dict:
return {
"i": idx, "score": 5, "pct": 50,
"exp": "Unknown", "match": [], "miss": [],
"note": "Assessment unavailable.", "kw": [], "pri": "Medium",
}
def _merge(job_dict: dict, quick: dict, detail: dict | None) -> dict:
score = int(quick.get("score", 5))
if detail:
score = int(detail.get("score", score))
return {
**job_dict,
"relevance_score": score,
"skills_match_percentage": int(detail.get("pct", 50)) if detail else 50,
"experience_match": detail.get("exp", "Unknown") if detail else "Unknown",
"matching_skills": ", ".join(detail.get("match", []) if detail else []),
"missing_skills": ", ".join(detail.get("miss", []) if detail else []),
"key_strengths": "",
"recommendation": detail.get("note", "") if detail else "",
"ats_keywords": ", ".join(detail.get("kw", []) if detail else []),
"application_priority": detail.get("pri", _score_to_priority(score)) if detail else _score_to_priority(score),
"_raw_assessment": detail or quick,
"_assessed_by": "",
}
def _score_to_priority(score: int) -> str:
if score >= 8: return "High"
if score >= 6: return "Medium"
return "Low"
# ── ModelPool ─────────────────────────────────────────────────────────────────
class ModelPool:
def __init__(self, model_configs: list[dict]):
self.workers_phase1 = [] # Fast models — Phase 1 quick scoring
self.workers_phase2 = [] # All models — Phase 2 detailed
for cfg in model_configs:
if not cfg.get("api_key"):
continue
try:
w = ModelWorker(
name=cfg["name"],
model=cfg["model"],
api_key=cfg["api_key"],
base_url=cfg["base_url"],
extra_body=cfg.get("extra_body", {}),
)
if cfg.get("phase1", True):
self.workers_phase1.append(w)
if cfg.get("phase2", True):
self.workers_phase2.append(w)
except Exception:
pass
if not self.workers_phase1 and not self.workers_phase2:
raise ValueError("No valid model configs provided.")
# If no phase1 workers, fall back to phase2
if not self.workers_phase1:
self.workers_phase1 = self.workers_phase2
print(f"{Fore.CYAN}Phase 1 workers: {', '.join(w.name for w in self.workers_phase1)}{Style.RESET_ALL}")
print(f"{Fore.CYAN}Phase 2 workers: {', '.join(w.name for w in self.workers_phase2)}{Style.RESET_ALL}")
def process_all(
self,
jobs: list,
compact_profile: str,
batch_size: int = DETAILED_BATCH_SIZE,
progress_callback=None,
) -> list[dict]:
"""
Detailed LLM assessment of the provided jobs.
Called by JobAssessor after keyword pre-scoring — only top jobs are passed.
Each worker grabs batches from shared queue the moment it finishes.
"""
job_dicts = [j.to_dict() if hasattr(j, "to_dict") else j for j in jobs]
total = len(job_dicts)
if total == 0:
return []
p2_names = ', '.join(w.name for w in self.workers_phase2)
print(f"\n{Fore.YELLOW}LLM assessing {total} jobs in batches of {batch_size} [{p2_names}]...{Style.RESET_ALL}")
print(f"{Fore.CYAN}Each worker grabs the next batch the moment it finishes.{Style.RESET_ALL}\n")
raw_details = self._run_pool(
job_dicts, compact_profile,
batch_fn=lambda w, batch, prof: w.detailed_batch(batch, prof),
batch_size=batch_size,
phase_name="LLM Detail",
workers=self.workers_phase2,
)
if progress_callback:
progress_callback(100, 100)
# Map flat detail list back to jobs by position
# Each batch of N jobs produces N detail dicts with "i" = 1..N
results = []
detail_idx = 0
for batch_start in range(0, total, batch_size):
batch_jobs = job_dicts[batch_start: batch_start + batch_size]
batch_size_actual = len(batch_jobs)
for pos, job in enumerate(batch_jobs):
# Find detail for this job (i = pos+1 within the batch)
detail = next(
(d for d in raw_details
if isinstance(d, dict)
and d.get("i") == pos + 1
and d not in results), # avoid reuse
None,
)
quick = {"job_index": batch_start + pos + 1, "score": job.get("_keyword_score", 5)}
results.append(_merge(job, quick, detail))
# Print worker stats
print(f"\n{Fore.CYAN}Worker stats:{Style.RESET_ALL}")
all_workers = list({w.name: w for w in self.workers_phase1 + self.workers_phase2}.values())
for w in all_workers:
print(f" {w.name}: {w.calls_done} calls, {w.jobs_done} jobs, {w.errors} errors")
results.sort(key=lambda x: x.get("relevance_score", 0), reverse=True)
return results
def _run_pool(
self,
job_dicts: list[dict],
compact_profile: str,
batch_fn,
batch_size: int,
phase_name: str,
workers: list = None,
) -> list:
if workers is None:
workers = self.workers_phase1
"""Generic worker-pool runner. Returns flat list of assessment dicts."""
batches = [job_dicts[i:i + batch_size] for i in range(0, len(job_dicts), batch_size)]
total_batches = len(batches)
bq: queue.Queue = queue.Queue()
for i, b in enumerate(batches):
bq.put((i, b, 0)) # (idx, batch, retry_count)
results: dict[int, list] = {}
lock = threading.Lock()
done_ctr = [0]
def worker_fn(worker: ModelWorker):
while True:
try:
batch_idx, batch, retries = bq.get(timeout=3)
except queue.Empty:
break
t0 = time.time()
try:
batch_result = batch_fn(worker, batch, compact_profile)
elapsed = time.time() - t0
with lock:
results[batch_idx] = batch_result
done_ctr[0] += 1
done = done_ctr[0]
pct = done / total_batches * 100
print(
f" [{phase_name}] {worker.name} · batch {batch_idx+1}/{total_batches} "
f"({len(batch)} jobs, {elapsed:.0f}s) — {pct:.0f}% done"
)
except RuntimeError as e:
# Server error — skip this model, put batch back for another worker
worker.errors += 1
if retries < 1:
bq.put((batch_idx, batch, retries + 1))
else:
with lock:
results[batch_idx] = [_neutral_detailed(i + 1) for i in range(len(batch))]
print(f" {Fore.RED}[{worker.name}] server error, re-queuing batch {batch_idx+1}{Style.RESET_ALL}")
except Exception as e:
worker.errors += 1
if retries < 2:
bq.put((batch_idx, batch, retries + 1))
else:
with lock:
results[batch_idx] = [_neutral_detailed(i + 1) for i in range(len(batch))]
print(f" {Fore.YELLOW}[{worker.name}] error batch {batch_idx+1}: {str(e)[:60]}{Style.RESET_ALL}")
finally:
bq.task_done()
threads = [threading.Thread(target=worker_fn, args=(w,), daemon=True) for w in workers]
for t in threads: t.start()
for t in threads: t.join()
# Flatten in order
flat = []
for idx in sorted(results.keys()):
items = results.get(idx, [])
if isinstance(items, list):
flat.extend(items)
return flat