diff --git a/api/src/stampy_chat/chat.py b/api/src/stampy_chat/chat.py index ca0b192..d70159e 100644 --- a/api/src/stampy_chat/chat.py +++ b/api/src/stampy_chat/chat.py @@ -1,8 +1,7 @@ from typing import Any, Callable, Dict, List -from langchain.chains import LLMChain, OpenAIModerationChain +from langchain.chains import LLMChain, OpenAIModerationChain, moderation from langchain.chat_models import ChatOpenAI -from langchain.embeddings.openai import OpenAIEmbeddings from langchain.memory import ChatMessageHistory, ConversationSummaryBufferMemory from langchain.prompts import ( BaseChatPromptTemplate, @@ -12,17 +11,13 @@ from langchain.prompts import ( ) from langchain.pydantic_v1 import Extra from langchain.schema import BaseMessage, ChatMessage, PromptValue, SystemMessage -from langchain.vectorstores import Pinecone -from stampy_chat.env import OPENAI_API_KEY, PINECONE_INDEX, PINECONE_NAMESPACE +from stampy_chat.env import OPENAI_API_KEY from stampy_chat.settings import Settings from stampy_chat.callbacks import StampyCallbackHandler, BroadcastCallbackHandler, LoggerCallbackHandler from stampy_chat.followups import StampyChain from stampy_chat.citations import make_example_selector -embeddings = OpenAIEmbeddings() -vectorstore = Pinecone(PINECONE_INDEX, embeddings.embed_query, "hash_id", namespace=PINECONE_NAMESPACE) - class ModerationError(ValueError): pass @@ -117,7 +112,12 @@ class LimitedConversationSummaryBufferMemory(ConversationSummaryBufferMemory): class ModeratedChatPrompt(ChatPromptTemplate): """Wraps a prompt with an OpenAI moderation check which will raise an exception if fails.""" - moderation_chain: OpenAIModerationChain = OpenAIModerationChain(error=True) + moderation_chain: OpenAIModerationChain = None + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + if not self.moderation_chain: + self.moderation_chain = OpenAIModerationChain(error=True, openai_api_key=OPENAI_API_KEY) def format_prompt(self, **kwargs: Any) -> PromptValue: """Raise an exception if the prompt is flagged as offensive by OpenAI.""" diff --git a/api/src/stampy_chat/citations.py b/api/src/stampy_chat/citations.py index bfc2141..3b8d706 100644 --- a/api/src/stampy_chat/citations.py +++ b/api/src/stampy_chat/citations.py @@ -8,14 +8,10 @@ from langchain.prompts import ( from langchain.pydantic_v1 import Extra from langchain.vectorstores import Pinecone -from stampy_chat.env import PINECONE_INDEX, PINECONE_NAMESPACE +from stampy_chat.env import PINECONE_INDEX, PINECONE_NAMESPACE, OPENAI_API_KEY from stampy_chat.callbacks import StampyCallbackHandler -embeddings = OpenAIEmbeddings() -vectorstore = Pinecone(PINECONE_INDEX, embeddings.embed_query, "hash_id", namespace=PINECONE_NAMESPACE) - - class ReferencesSelector(SemanticSimilarityExampleSelector): """Get examples with enumerated indexes added.""" @@ -65,6 +61,8 @@ class ReferencesSelector(SemanticSimilarityExampleSelector): def make_example_selector(k: int, **params) -> ReferencesSelector: + embeddings = OpenAIEmbeddings(openai_api_key=OPENAI_API_KEY) + vectorstore = Pinecone(PINECONE_INDEX, embeddings.embed_query, "hash_id", namespace=PINECONE_NAMESPACE) return ReferencesSelector(vectorstore=vectorstore, **params) diff --git a/api/src/stampy_chat/env.py b/api/src/stampy_chat/env.py index 51e31ec..c0df97d 100644 --- a/api/src/stampy_chat/env.py +++ b/api/src/stampy_chat/env.py @@ -1,5 +1,5 @@ import os -import openai +# import openai import pinecone if os.path.exists('.env'): @@ -16,7 +16,6 @@ DISCORD_LOGGING_URL = os.environ.get('LOGGING_URL') ### OpenAI ### OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY') -openai.api_key = OPENAI_API_KEY # non-optional ### Models ### EMBEDDING_MODEL = os.environ.get("EMBEDDING_MODEL", "text-embedding-ada-002")