handle missing OpenAI key

This commit is contained in:
Daniel O'Connell
2023-10-15 21:57:06 +02:00
parent 7440021964
commit a04adf68a4
3 changed files with 12 additions and 15 deletions
+8 -8
View File
@@ -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."""
+3 -5
View File
@@ -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)
+1 -2
View File
@@ -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")