47 lines
960 B
Python
47 lines
960 B
Python
from celery import shared_task
|
|
from .models import Track
|
|
|
|
import dotenv
|
|
|
|
import os
|
|
import requests
|
|
|
|
import time
|
|
import json
|
|
from random import choice
|
|
|
|
dotenv.load_dotenv()
|
|
|
|
UPDATE_FREQ = 10 # How often to save to database
|
|
|
|
@shared_task
|
|
def get_banter(id):
|
|
track = Track.objects.get(pk=id)
|
|
r = requests.post(
|
|
os.getenv('AI_ENDPOINT'),
|
|
json={
|
|
'model': os.getenv('AI_MODEL'),
|
|
'prompt': str(track)
|
|
},
|
|
stream=True
|
|
)
|
|
r.raise_for_status()
|
|
|
|
token = 0
|
|
|
|
for line in r.iter_lines():
|
|
body = json.loads(line)
|
|
response_part = body.get('response', '')
|
|
token += 1
|
|
track.banter += response_part
|
|
|
|
if 'error' in body:
|
|
raise Exception(body['error'])
|
|
|
|
if body.get('done', False):
|
|
track.banter_done = True
|
|
track.save()
|
|
return body['context']
|
|
elif token % UPDATE_FREQ == 0:
|
|
track.save()
|