aegislm / backend /queue /producer.py
ACA050's picture
Upload 50 files
1a4aa87 verified
Raw
History Blame Contribute Delete
8.76 kB
"""
Job Producer for Evaluation Queue
Provides job submission functionality for the evaluation queue.
Handles job creation, validation, and queueing.
"""
import hashlib
import json
import uuid
from datetime import datetime
from typing import Optional
from backend.core.config import settings
from backend.logging.logger import get_logger
from .job_schema import (
JobStatus,
JobType,
JobPriority,
EvaluationJob,
JobSubmissionRequest,
)
from .status_tracker import get_status_tracker
logger = get_logger("queue.producer", component="queue")
class JobProducer:
"""
Produces and submits evaluation jobs to the queue.
Responsibilities:
- Validate job submissions
- Generate config hashes for reproducibility
- Create job records in database
- Submit jobs to the queue
"""
def __init__(self):
self._status_tracker = get_status_tracker()
def _generate_config_hash(
self,
model_name: str,
model_version: str,
dataset_name: str,
dataset_version: str,
mutation_depth: int,
attack_types: list[str],
) -> str:
"""Generate SHA256 hash of configuration for reproducibility."""
config = {
"model_name": model_name,
"model_version": model_version,
"dataset_name": dataset_name,
"dataset_version": dataset_version,
"mutation_depth": mutation_depth,
"attack_types": sorted(attack_types),
}
config_str = json.dumps(config, sort_keys=True)
return hashlib.sha256(config_str.encode()).hexdigest()
async def submit_job(
self,
request: JobSubmissionRequest,
submitted_by: Optional[str] = None,
) -> EvaluationJob:
"""
Submit a new evaluation job.
Args:
request: Job submission request
submitted_by: API key owner who submitted the job
Returns:
Created evaluation job
"""
try:
# Generate config hash
config_hash = self._generate_config_hash(
model_name=request.model_name,
model_version=request.model_version,
dataset_name=request.dataset_name,
dataset_version=request.dataset_version,
mutation_depth=request.mutation_depth,
attack_types=request.attack_types,
)
# Create job
job = EvaluationJob(
job_id=uuid.uuid4(),
job_type=request.job_type,
model_name=request.model_name,
model_version=request.model_version,
dataset_name=request.dataset_name,
dataset_version=request.dataset_version,
config_hash=config_hash,
priority=request.priority,
submitted_by=submitted_by,
status=JobStatus.PENDING,
progress=0.0,
total_samples=0, # Will be updated when job starts
completed_samples=0,
failed_samples=0,
checkpoint_interval=request.checkpoint_interval,
created_at=datetime.utcnow(),
metadata={
"mutation_depth": request.mutation_depth,
"attack_types": request.attack_types,
"max_concurrency": request.max_concurrency,
"sampling_config": request.sampling_config,
},
)
# Create job in database
await self._status_tracker.create_job(job)
# Update status to queued
await self._status_tracker.update_job_status(
job.job_id,
JobStatus.QUEUED,
)
job.status = JobStatus.QUEUED
job.queued_at = datetime.utcnow()
# Add to in-memory queue (for now, using simple list)
# In production, this would use Redis/RQ/Celery
_job_queue.append(job)
logger.info(
"Job submitted",
job_id=str(job.job_id),
job_type=job.job_type,
model=job.model_name,
dataset=job.dataset_name,
priority=job.priority,
)
return job
except Exception as e:
logger.error(
"Failed to submit job",
error=str(e),
)
raise
async def submit_benchmark_job(
self,
model_name: str,
model_version: str,
dataset_version: str,
submitted_by: Optional[str] = None,
priority: JobPriority = JobPriority.NORMAL,
mutation_depth: int = 2,
attack_types: Optional[list[str]] = None,
max_concurrency: int = 4,
checkpoint_interval: int = 10,
) -> EvaluationJob:
"""
Submit a benchmark evaluation job.
Convenience method for submitting standard benchmark jobs.
Args:
model_name: Name of the model to evaluate
model_version: Model version
dataset_version: Dataset version to use
submitted_by: API key owner
priority: Job priority
mutation_depth: Mutation depth
attack_types: List of attack types
max_concurrency: Maximum concurrent samples
checkpoint_interval: Checkpoint interval
Returns:
Created evaluation job
"""
request = JobSubmissionRequest(
job_type=JobType.BENCHMARK,
model_name=model_name,
model_version=model_version,
dataset_name="default",
dataset_version=dataset_version,
priority=priority,
mutation_depth=mutation_depth,
attack_types=attack_types or ["jailbreak"],
max_concurrency=max_concurrency,
checkpoint_interval=checkpoint_interval,
)
return await self.submit_job(request, submitted_by)
async def cancel_job(
self,
job_id: uuid.UUID,
) -> bool:
"""
Cancel a job.
Args:
job_id: The job ID to cancel
Returns:
True if cancelled, False if not found or already completed
"""
try:
job = self._status_tracker.get_cached_job(job_id)
if job is None:
logger.warning(
"Job not found for cancellation",
job_id=str(job_id),
)
return False
# Check if job can be cancelled
if job.status in [JobStatus.COMPLETED, JobStatus.FAILED, JobStatus.CANCELLED]:
logger.info(
"Job already in terminal state",
job_id=str(job_id),
status=job.status,
)
return False
# Update status
await self._status_tracker.update_job_status(
job_id,
JobStatus.CANCELLED,
)
logger.info(
"Job cancelled",
job_id=str(job_id),
)
return True
except Exception as e:
logger.error(
"Failed to cancel job",
job_id=str(job_id),
error=str(e),
)
return False
def get_queue_size(self) -> int:
"""Get current queue size."""
return len(_job_queue)
def get_pending_jobs(self) -> list[EvaluationJob]:
"""Get all pending jobs in queue."""
return [job for job in _job_queue if job.status == JobStatus.QUEUED]
# In-memory job queue (for demonstration)
# In production, this would be replaced with Redis/RQ/Celery
_job_queue: list[EvaluationJob] = []
# Global instance
_producer: Optional[JobProducer] = None
def get_job_producer() -> JobProducer:
"""Get the global job producer instance."""
global _producer
if _producer is None:
_producer = JobProducer()
return _producer
__all__ = [
"JobProducer",
"get_job_producer",
]