CinAI-ShotSense / utils.py
Ashish1722's picture
Upload 2 files
58ce4ca verified
Raw
History Blame Contribute Delete
5.62 kB
# utils.py
import os
import json
import requests
import time
from openai import AzureOpenAI
import streamlit as st
import tempfile, base64
import pymupdf as pf
from json_repair import repair_json
from pymongo import MongoClient
from dotenv import load_dotenv
# Load environment variables
load_dotenv()
# MongoDB setup
MONGO_URL = "mongodb+srv://ashish171222:zIWE6zGJ642nlpij@cluster0.hwivz3z.mongodb.net/mydatabase?retryWrites=true&w=majority"
mongoclient = MongoClient(MONGO_URL)
db = mongoclient['script']
collection = db['scripts']
# Azure OpenAI setup
client = AzureOpenAI(
azure_endpoint = 'https://cinai2.openai.azure.com/',
api_key= '660f1642218b4ceeb4e2553d24f0e0bd',
api_version="2024-05-01-preview"
)
def create_id(text):
list_chunk = []
thread = client.beta.threads.create()
batch_size = 5 # Number of pages per message
accumulated_text = ""
page_count = 0
for i, page in enumerate(text):
# Extract and encode page text
page_text = page.get_text().encode("utf8").decode("utf8")
# Accumulate text for batch
accumulated_text += f"Page {i + 1} Data:\n{page_text}\n\n"
page_count += 1
# If we reach the batch size or it's the last page, send the message
if page_count >= batch_size or i == len(text) - 1:
try:
message = client.beta.threads.messages.create(
thread_id=thread.id,
role="user",
content=accumulated_text
)
print(f"Successfully sent pages {i - page_count + 2} to {i + 1}.")
except Exception as e:
print(f"Failed to send message for pages {i - page_count + 2} to {i + 1}: {e}")
# Reset accumulators for the next batch
accumulated_text = ""
page_count = 0
# Combine chunks into larger messages
list_chunk.append(page_text)
thread_id = thread.id
st.write(thread_id)
return thread_id
def search_or_insert_data(name, text, path):
view_pdf(path)
with st.spinner('Sending file to the assistant'):
thread_id = create_id(text)
return thread_id
def load_script(script_file):
if 'name' not in st.session_state:
st.session_state['name'] = None
if st.session_state['name'] != script_file.name:
st.session_state['name'] = script_file.name
with tempfile.NamedTemporaryFile(delete=False) as tmp_file:
tmp_file.write(script_file.read())
temp_file_path = tmp_file.name
st.session_state['path'] = temp_file_path
st.session_state['script_text'] = pf.open(st.session_state['path'])
name = st.session_state['name']
if 'path' not in st.session_state:
with tempfile.NamedTemporaryFile(delete=False) as tmp_file:
tmp_file.write(script_file.read())
temp_file_path = tmp_file.name
st.session_state['path'] = temp_file_path
st.session_state['script_text'] = pf.open(st.session_state['path'])
if pf.open(st.session_state['path']):
st.session_state['script_text'] = pf.open(st.session_state['path'])
result = search_or_insert_data(name, st.session_state['script_text'], st.session_state['path'])
else:
st.session_state['path'] = temp_file_path
st.session_state['script_text'] = pf.open(st.session_state['path'])
result = search_or_insert_data(name, st.session_state['script_text'], st.session_state['path'])
return result
def view_pdf(script):
with open(script, "rb") as f:
base64_pdf = base64.b64encode(f.read()).decode('utf-8')
pdf_display = F'<embed src="data:application/pdf;base64,{base64_pdf}" width="525" height="750" type="application/pdf">'
st.sidebar.markdown(pdf_display, unsafe_allow_html=True)
def get_assistant_id(feature):
assistant_ids = {
"Script Analysis": "asst_0TVOqfDUPuaSxtea11xa7DB0",
"Franchise Continuity Checker": "asst_iirHFFSynUu1kBPRP8KT7wPl",
"Audience Reaction": "asst_kr92jSrWpbdEI9wSHl2OIOA2",
"Character Analysis": "asst_2xl7bqCuNlvfawVBCkSSRbIP",
"Crossover Potential": "asst_Kz39dY89bAF83rFCBhwTmvul",
"VFX Potential":"asst_WZPuGYRu6zU3XomrOietVAS5",
"Market Analysis": "asst_ykSNeNu74RsJPkOxLTPYHQ36"
}
return assistant_ids.get(feature)
def run_analysis(thread_id, feature, additional_context=None):
if additional_context:
user_message = {
"role": "user",
"content": json.dumps(additional_context)
}
client.beta.threads.messages.create(
thread_id=thread_id,
role="user",
content= json.dumps(additional_context)
)
run = client.beta.threads.runs.create(
thread_id=thread_id,
assistant_id=get_assistant_id(feature=feature)
)
while run.status in ['queued', 'in_progress', 'cancelling']:
time.sleep(1)
run = client.beta.threads.runs.retrieve(
thread_id=thread_id,
run_id=run.id
)
if run.status == 'completed':
time.sleep(15)
messages = client.beta.threads.messages.list(thread_id=thread_id)
analysis = next((msg.content[0].text.value for msg in reversed(list(messages)) if msg.role == "assistant"), "")
return analysis
else:
return f"Error: Run status is {run.status}"