Files
python_test/smartassist/src/backend.py
T

230 lines
8.6 KiB
Python

# Import the necessary functions from ollama, Flask, requests, threading
from ollama import Client
from flask import Flask, request, jsonify, send_from_directory, render_template, session, make_response
from flask_cors import CORS, cross_origin # CORS stands for Cross-Origin Resource Sharing. This is necessary to allow the frontend to make requests to our backend.
import requests
import json
import logging
import os
import utils
from utils import GlobalState
# Create a logger for this module
global_state = GlobalState() # Import the singleton that holds global states (e.g., logger)
logger = global_state.get_logger(__name__) # Logger for this module, inherit properties of the root logger
# Find out the path to current directory according to the Python interpreter (venv)
logger.debug("Current working directory: %s", os.getcwd())
# Initialize a Flask application
app = Flask(__name__)
app.config['STATIC_FOLDER'] = 'static' # Adjust if needed
# Increase the maximum cookie size
app.config['SESSION_COOKIE_SAMESITE'] = 'Lax'
app.config['SESSION_COOKIE_SIZE_LIMIT'] = 4096 * 2 # Allow up to 8KB cookies
# Set the secret key for session management
secret_key = os.urandom(24)
app.config['SECRET_KEY'] = secret_key # When do I need this. How is it retained between sessions?
# Optionally set other configuration options
app.config['SESSION_PERMANENT'] = False # Session will expire after each request
app.config['SESSION_TYPE'] = 'filesystem' # Store sessions on the filesystem
logger.debug("flask app template folder: %s", app.template_folder)
@app.route('/')
def index():
"""
This route serves index.html to connecting clients
"""
session['chat_history'] = [] # The session object (actually, a dictonary) holds the chat session
logger.debug("Entering route '/'")
api_endpoint = global_state.get_backend_api_ep() # Retrieve the environment variable
host_url = global_state.get_host_url()
use_model = global_state.get_llm()
logger.debug("Backend API endpoint:\t%s", api_endpoint)
logger.debug("Host of LLMs:\t\t%s", host_url)
logger.debug("LLM to use:\t\t\t%s", use_model)
with open('smartassist/src/html/client.html', 'r') as f:
client_html = f.read()
# logger.debug("Client HTML (first few characters): %s", client_html[:50]) # Print to see if it's loading
# logger.debug("Client HTML (all characters): %s", client_html) # Print to see if it's loading
return render_template('index.html', api_endpoint=api_endpoint, use_model = use_model, client_content=client_html)
@app.route('/set_session')
def set_session():
resp = make_response()
resp.set_cookie('session', 'some-value', samesite='None', secure=True) # Add SameSite attribute here
return resp
# @app.route('/profile')
# def profile():
# # Retrieve data from the session
# user_id = session.get('user_id')
# if user_id:
# return f'User ID: {user_id}'
# else:
# return 'No user ID found'
@app.route('/<path:filename>')
def serve_static(filename):
return send_from_directory(app.config['STATIC_FOLDER'], filename)
# CORS(app, resources={
# r"/api/chat": {
# "origins": "*",
# "headers": ["Origin", "Content-Type", "Authorization"],
# }
# })
CORS(app, resources={
r"/api/chat": {
"origins": "*"
}
})
@app.route('/api/tags', methods=['GET'])
def tag(url = "http://localhost:11434/api/tags", headers = None):
"""Get a list of models for the server located at url."""
try:
logger.debug(f"url: {url} headers: {headers}")
response = requests.get(url, headers=headers)
return response.json()
# return response
except requests.exceptions.RequestException as e:
logger.error("Request Exception: %s", str(e))
return {'error': 'Failed to process request'}
@app.route('/api/chat', methods=['POST'])
def chat():
# def chat(model = "phi3:mini"):
"""
This function handles the chat. The frontend client (web browser) calls the
backend server through this endpoint (/api/chat) that manage queries
to the LLM (Large Language Model) server and it also manages the response
from the LLM server.
"""
# Get the message from the JSON in the request body
data = request.get_json()
message = data.get('query')
url_server = data.get('url_server', global_state.get_host_url()) # Use provided URL or current if not provided
# url_server = data.get('url_server', "https://ollama-test.wara-ops.org/api/generate") # Use provided URL or default
model = data.get('model', global_state.get_llm()) # Use provided model or current if not provided
# Get chat history from session storage (e.g., a dictionary)
chat_history = session.get('chat_history', [])
# Add the new message to the chat history
chat_history.append({'role': 'user', 'message': message})
# Update the session with the new chat history
session['chat_history'] = chat_history
# Create the data dictionary with chat history
data_to_send = {
"model": model,
'prompt': '\n'.join([f"{item['role']}: {item['message']}" for item in chat_history]),
"stream": False
}
url = url_server
# TODO: This section should only run when changing to new endpoint...
# begin refactor ###################################
headers = { # Set default header
"Content-Type": "application/json",
}
found_endpoint = False
endpoints = global_state.get_endpoints()
for endpoint in endpoints:
if endpoint["url"] == url: # Look for endpoint with this url
found_endpoint = True
if endpoint["provider"] == "ollama": # Currently only supporting ollama servers
if "requestOptions" in endpoint: # Check if authentication is needed
headers.update({
"Authorization": endpoint["requestOptions"]["headers"]["Authorization"]
})
if found_endpoint == False:
# Raise some error or whatever...
logger.debug(f"Host {url} not found")
# end refactor ###################################
try:
url = url + "/api/generate"
logger.debug(f"url: {url} headers: {headers}")
response = requests.post(url,
headers=headers,
data=json.dumps(data_to_send))
response.raise_for_status() # Raise an exception for bad status codes
llm_response = response.json()['response'] # Assuming the LLM's response is under 'response' key
chat_history.append({'role': 'assistant', 'message': llm_response}) # Add assistant response to chat history
logger.debug(f"Chat History: {chat_history}")
return response.json()
except requests.exceptions.RequestException as e:
logger.error("Request Exception: %s", str(e))
return jsonify({'error': 'Failed to process request'}), 500
except json.JSONDecodeError as e:
logger.error("JSON Decode Error: %s", str(e)) # Corresponds to print(f"JSON Decode Error: {e}")
return jsonify({'error': 'Invalid JSON response from server'}), 500
@app.route('/api/endpoints', methods=['GET'])
def get_endpoints():
# Replace this with your actual logic to fetch endpoint data
endpoints = [
{'endpoint': 'endpoint1', 'llm': 'llm1'},
{'endpoint': 'endpoint2', 'llm': 'llm2'},
{'endpoint': 'endpoint3', 'llm': 'llm3'},
{'endpoint': 'endpoint4', 'llm': 'llm4'},
]
return jsonify(endpoints)
@app.route('/smartassist', methods=["POST"])
def smartassist():
# Extract the query from the incoming JSON data
data = request.json
user_query = data['query']
# Get the response from the OLLAMA API based on the user's query
# NOTE: Should we append message history here? Maybe interact with SQLlite?
response = get_response(user_query)
# Return the response as a JSON object in the HTTP response
return jsonify({"response": response})
def get_response(user_query):
client = Client() # Create a client object for interacting with OLLAMA API
response = client.generate_response(user_query) # Generate and retrieve the response based on user's query
return response
def run_flask(fport=5005):
"""
Starts the Flask server
"""
# Flask endpoint for user interaction
logger.debug("Entering run_flask()")
# app.run(port = str(str(fport)), debug=False)
app.run(port = str(str(fport)), debug=True)
# app.run(port=5000, debug=True, use_reloader=False)
logger.debug("Exiting run_flask()")
if __name__ == '__main__':
# Run the Flask application
run_flask()