66 lines
2.0 KiB
Python
66 lines
2.0 KiB
Python
# All comments are in English.
|
|
|
|
from celery_app import celery
|
|
import requests
|
|
import json
|
|
import uuid
|
|
from models import Message, Conversation
|
|
from sqlalchemy.orm import sessionmaker
|
|
from sqlalchemy import create_engine
|
|
from config import Config
|
|
|
|
engine = create_engine(Config.DATABASE_URL)
|
|
SessionLocal = sessionmaker(bind=engine)
|
|
|
|
@celery.task(bind=True, max_retries=3)
|
|
def generate_task(self, conversation_id, model, prompt):
|
|
"""Runs an Ollama generation asynchronously."""
|
|
try:
|
|
url = f"{Config.OLLAMA_BASE_URL}/api/generate"
|
|
payload = {"model": model, "prompt": prompt}
|
|
|
|
r = requests.post(url, json=payload)
|
|
r.raise_for_status()
|
|
output = r.json().get("response", "")
|
|
|
|
db = SessionLocal()
|
|
conv = db.query(Conversation).filter_by(id=cid, user_email=email).first()
|
|
if not conv:
|
|
logger.error("conversation_access_denied")
|
|
return None
|
|
|
|
msg = Message(
|
|
id=str(uuid.uuid4()),
|
|
conversation_id=conversation_id,
|
|
role="assistant",
|
|
content=output,
|
|
)
|
|
db.add(msg)
|
|
db.commit()
|
|
db.close()
|
|
|
|
return output
|
|
|
|
except Exception as e:
|
|
raise self.retry(exc=e, countdown=2)
|
|
|
|
|
|
|
|
@celery.task(name="tasks.system_admin_generate", bind=True)
|
|
def system_admin_generate(self, email, cid, model, prompt):
|
|
logger.info("system_admin_task_started")
|
|
conv = db.query(Conversation).filter_by(id=cid, user_email=email).first()
|
|
return run_ollama(cid, model, prompt)
|
|
|
|
@celery.task(name="tasks.corp_admin_generate", bind=True)
|
|
def corp_admin_generate(self, email, cid, model, prompt):
|
|
logger.info("corp_admin_task_started")
|
|
conv = db.query(Conversation).filter_by(id=cid, user_email=email).first()
|
|
return run_ollama(cid, model, prompt)
|
|
|
|
@celery.task(name="tasks.user_generate", bind=True)
|
|
def user_generate(self, email, cid, model, prompt):
|
|
logger.info("user_task_started")
|
|
conv = db.query(Conversation).filter_by(id=cid, user_email=email).first()
|
|
return run_ollama(cid, model, prompt)
|