import os from functools import lru_cache from langchain.chat_models import AzureChatOpenAI, ChatOpenAI from langchain.chat_models.base import BaseChatModel from langchain.embeddings import OpenAIEmbeddings from langchain.schema import Document from langchain.vectorstores import Qdrant, VectorStore from edu_assistant.utils.qdrant_utils import load_qdrant_client @lru_cache(maxsize=1) def load_llm() -> BaseChatModel: if os.environ.get("AZURE_OPENAI"): llm = AzureChatOpenAI( openai_api_type="azure", openai_api_key=os.environ.get("AZURE_OPENAI_API_KEY"), openai_api_base=os.environ.get("AZURE_OPENAI_API_BASE"), openai_api_version="2023-05-15", deployment_name=os.environ.get("AZURE_OPENAI_DEPLOYMENT_ID", "gpt-35-turbo"), model="gpt-3.5-turbo", temperature=0, ) else: llm = ChatOpenAI( openai_api_key=os.environ.get("OPENAI_API_KEY"), openai_proxy=os.environ.get("OPENAI_PROXY", ""), model="gpt-3.5-turbo", ) return llm @lru_cache(maxsize=1) def load_gpt4_llm() -> BaseChatModel: llm = ChatOpenAI( openai_api_key=os.environ.get("OPENAI_API_KEY"), openai_proxy=os.environ.get("OPENAI_PROXY", ""), model="gpt-4", ) return llm @lru_cache(maxsize=1) def load_gpt4_flag() -> bool: return os.environ.get("CODEDOG_ENABLE_GPT4") is not None @lru_cache(maxsize=1) def load_embeddings(): if os.environ.get("AZURE_OPENAI"): embeddings = OpenAIEmbeddings( openai_api_type="azure", openai_api_key=os.environ.get("AZURE_OPENAI_API_KEY"), openai_api_base=os.environ.get("AZURE_OPENAI_API_BASE"), openai_api_version="2023-05-15", deployment=os.environ.get("AZURE_OPENAI_EMBEDDING_DEP_ID", ""), ) else: embeddings = OpenAIEmbeddings( openai_api_key=os.environ.get("OPENAI_API_KEY"), openai_proxy=os.environ.get("OPENAI_PROXY", ""), ) return embeddings @lru_cache(maxsize=10) def load_vectorstore(collection_name: str = "default") -> VectorStore: if os.environ.get("QDRANT_API"): client = load_qdrant_client() embeddings = load_embeddings() doc_store = Qdrant(client=client, collection_name=collection_name, embeddings=embeddings) return doc_store return None @lru_cache(maxsize=20) def escape_for_prompt(text: str) -> str: """escape cruly brackets in text for generate prompt. Args: text (str): input string. Returns: str: escaped string. """ return text.replace("{", "{{").replace("}", "}}") def shrink_docs(docs: list[Document], max_size=50): """shrink source docs content size for display. Args: docs (dict): Retrieval Chain returned docs. """ for doc in docs: doc.page_content = doc.page_content[:max_size] + ".." return docs