Skip to content

Commit 26adb7d

Browse files
Merge pull request #1 from akshaypardhanani/feat/add-support-for-openrouter
feat: Add support for OpenRouter
2 parents ddd5f7a + 9306598 commit 26adb7d

12 files changed

Lines changed: 1250 additions & 29 deletions

File tree

.gitignore

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -46,3 +46,11 @@ MANIFEST
4646
*.pyc
4747
__pycache__/
4848
*~
49+
50+
digest.txt
51+
uv.lock
52+
.vscode/launch.json
53+
amadeusgpt/modules_embedding.pickle
54+
temp_answer.json
55+
56+
logs/

Makefile

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
11
export streamlit_app=True
22
app:
33

4-
streamlit run amadeusgpt/app.py --server.fileWatcherType none --server.maxUploadSize 1000
4+
uv run --with streamlit streamlit run amadeusgpt/app.py --server.fileWatcherType none --server.maxUploadSize 1000

amadeusgpt/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44
SOURCE CODE: https://github.com/AdaptiveMotorControlLab/AmadeusGPT
55
Apache-2.0 license
66
"""
7+
import streamlit as st
78

89
from matplotlib import pyplot as plt
910

amadeusgpt/analysis_objects/llm.py

Lines changed: 13 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313

1414
from amadeusgpt.programs.sandbox import Sandbox
1515
from amadeusgpt.utils import AmadeusLogger, QA_Message, create_qa_message
16+
from amadeusgpt.utils.openai_adapter import OpenAIAdapter
1617

1718
from .base import AnalysisObject
1819

@@ -22,6 +23,7 @@ class LLM(AnalysisObject):
2223
prices = {
2324
"gpt-4o": {"input": 5 / 10**6, "output": 15 / 10**6},
2425
"gpt-4o-mini": {"input": 0.15 / 10**6, "output": 0.6 / 10**6},
26+
"thudm/glm-z1-32b:free": {"input": 0, "output": 0},
2527
}
2628
total_cost = 0
2729

@@ -65,17 +67,17 @@ def connect_gpt_oai_1(self, messages, **kwargs):
6567
This is routed to openai > 1.0 interfaces
6668
"""
6769

68-
if self.config.get("use_streamlit", False):
69-
if "OPENAI_API_KEY" in os.environ:
70-
openai.api_key = os.environ["OPENAI_API_KEY"]
71-
else:
72-
openai.api_key = os.environ["OPENAI_API_KEY"]
70+
# if self.config.get("use_streamlit", False):
71+
# if "OPENAI_API_KEY" in os.environ:
72+
# openai.api_key = os.environ["OPENAI_API_KEY"]
73+
# else:
74+
# openai.api_key = os.environ["OPENAI_API_KEY"]
7375
response = None
7476
# gpt_model is default to be the cls.gpt_model, which can be easily set
75-
gpt_model = self.gpt_model
77+
# gpt_model = self.gpt_model
7678
# in streamlit app, "gpt_model" is set by the text box
7779

78-
client = OpenAI()
80+
client = OpenAIAdapter().get_client()
7981

8082
if self.config.get("use_streamlit", False):
8183
if "gpt_model" in st.session_state:
@@ -89,7 +91,7 @@ def connect_gpt_oai_1(self, messages, **kwargs):
8991

9092
# the usage was recorded from the last run. However, since we have many LLMs that
9193
# share the call of this function, we will need to store usage and retrieve them from the database class
92-
num_retries = 3
94+
num_retries = 1
9395
for _ in range(num_retries):
9496
try:
9597
json_data = {
@@ -259,7 +261,7 @@ def speak(self, sandbox: Sandbox, image: np.ndarray):
259261
multi_image_content=multi_image_content,
260262
in_place=True,
261263
)
262-
response = self.connect_gpt(self.context_window, max_tokens=2000)
264+
response = self.connect_gpt(self.context_window, max_tokens=20000)
263265
text = response.choices[0].message.content.strip()
264266

265267
print("description of the image frame provided")
@@ -319,7 +321,7 @@ def speak(
319321

320322
self.update_history("user", query)
321323

322-
response = self.connect_gpt(self.context_window, max_tokens=2000)
324+
response = self.connect_gpt(self.context_window, max_tokens=20000)
323325
text = response.choices[0].message.content.strip()
324326
# need to keep the memory of the answers from LLM
325327
self.update_history("assistant", text)
@@ -374,7 +376,7 @@ def speak(self, qa_message):
374376
Can you correct the code? Make sure you only write one function which is the updated function.
375377
"""
376378
self.update_history("user", query)
377-
response = self.connect_gpt(self.context_window, max_tokens=4096)
379+
response = self.connect_gpt(self.context_window, max_tokens=20000)
378380
text = response.choices[0].message.content.strip()
379381
print(text)
380382
pattern = r"```python(.*?)```"

amadeusgpt/app.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,7 @@
11
import os
22
import traceback
33

4-
import streamlit as st
5-
6-
from amadeusgpt import app_utils
4+
from amadeusgpt import app_utils, st
75
from amadeusgpt.utils import validate_openai_api_key
86

97
# Set page configuration
@@ -21,6 +19,8 @@ def main():
2119
if "exist_valid_openai_api_key" not in st.session_state:
2220
if "OPENAI_API_KEY" in os.environ:
2321
st.session_state["exist_valid_openai_api_key"] = True
22+
elif "OPENROUTER_API_KEY" in os.environ:
23+
st.session_state["exist_valid_openai_api_key"] = True
2424
else:
2525
st.session_state["exist_valid_openai_api_key"] = False
2626

@@ -30,6 +30,8 @@ def valid_api_key():
3030
print("inside valid api key function")
3131
if "OPENAI_API_KEY" in os.environ:
3232
api_token = os.environ["OPENAI_API_KEY"]
33+
elif "OPENROUTER_API_KEY" in os.environ:
34+
api_token = os.environ["OPENROUTER_API_KEY"]
3335
else:
3436
api_token = st.session_state["openAI_token"]
3537
check_valid = validate_openai_api_key(api_token)

amadeusgpt/configs/Horse_template.yaml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ keypoint_info:
77
nose: "nose"
88
neck: "neck"
99
llm_info:
10+
gpt_model: "thudm/glm-z1-32b:free"
1011
keep_last_n_messages: 2
1112
object_info:
1213
load_objects_from_disk: false

amadeusgpt/integration_module_hub.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,18 @@
11
import os
22
import pickle
33

4-
from openai import OpenAI
4+
from amadeusgpt.utils.openai_adapter import OpenAIAdapter
55
from sklearn.metrics.pairwise import cosine_similarity
66

7+
from amadeusgpt import st
78
from amadeusgpt.programs.api_registry import INTEGRATION_API_REGISTRY
89

9-
client = OpenAI()
10-
1110

1211
class IntegrationModuleHub:
1312
def __init__(self):
1413
self.amadeus_root = os.path.dirname(os.path.realpath(__file__))
14+
self.client = OpenAIAdapter(
15+
api_key=st.session_state.get("OPENAI_API_KEY") or st.session_state.get("OPENROUTER_API_KEY") or st.session_state.get("openAI_token")).get_client()
1516

1617
def save_embeddings(self):
1718
result = {}
@@ -20,7 +21,7 @@ def save_embeddings(self):
2021
docstring = module_info["description"]
2122
text = docstring.replace("\n", " ")
2223
embedding = (
23-
client.embeddings.create(input=[text], model=model).data[0].embedding
24+
self.client.embeddings.create(input=[text], model=model).data[0].embedding
2425
)
2526
result[module_name] = embedding
2627
if len(result) > 0:
@@ -35,7 +36,7 @@ def match_module(self, query):
3536
model = "text-embedding-3-small"
3637

3738
query_embedding = (
38-
client.embeddings.create(input=[query], model=model).data[0].embedding
39+
self.client.embeddings.create(input=[query], model=model).data[0].embedding
3940
)
4041

4142
if not os.path.exists(

amadeusgpt/utils/__init__.py

Lines changed: 3 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
from amadeusgpt.analysis_objects.event import Event
1010
from amadeusgpt.logger import AmadeusLogger
1111
from IPython.display import Markdown, Video, display, HTML
12+
from amadeusgpt.utils.openai_adapter import OpenAIAdapter
1213

1314
def filter_kwargs_for_function(func, kwargs):
1415
sig = inspect.signature(func)
@@ -36,13 +37,8 @@ def parse_error_message_from_python():
3637
return traceback_str
3738

3839
def validate_openai_api_key(key):
39-
import openai
40-
openai.api_key = key
41-
try:
42-
openai.models.list()
43-
return True
44-
except openai.AuthenticationError:
45-
return False
40+
client = OpenAIAdapter(key)
41+
return client.validate()
4642

4743
def flatten_tuple(t):
4844
"""
Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
from typing import Union
2+
from sentence_transformers import SentenceTransformer
3+
4+
class EmbeddingModel:
5+
def __init__(self, provider_type: str, model_name: str = "BAAI/bge-m3", openai_model: str = "text-embedding-3-small"):
6+
self.provider_type = provider_type
7+
if self.provider_type == "openrouter":
8+
self.embedding_model = SentenceTransformer(model_name)
9+
self.embeddings = self._Embeddings(self.provider_type, self.embedding_model or openai_model)
10+
11+
class _Embeddings:
12+
def __init__(self, provider: str, model: Union[str, SentenceTransformer]):
13+
self.provider = provider
14+
self.model = model
15+
16+
def create(self, input: list[str], model: Union[str, SentenceTransformer]):
17+
if self.provider == "openrouter":
18+
vecs = self.model.encode(input, normalize_embeddings=True, convert_to_tensor=False)
19+
if isinstance(input, str):
20+
vecs = [vecs] # (1, D)
21+
22+
# Convert each embedding to a plain Python list
23+
vecs = [v.tolist() for v in vecs]
24+
25+
# Mimic OpenAI's Embedding.create() response
26+
return type("FakeResponse", (), {
27+
"data": [
28+
type("EmbeddingData", (), {
29+
"embedding": emb,
30+
"index": idx
31+
})()
32+
for idx, emb in enumerate(vecs)
33+
]
34+
})()
35+
else:
36+
return self.model.embeddings.create(input=input, model=model)

amadeusgpt/utils/openai_adapter.py

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
import openai
2+
from amadeusgpt.utils.embedding_adapter import EmbeddingModel
3+
4+
class OpenAIAdapter:
5+
_instance = None
6+
_initialized = False
7+
8+
def __new__(cls, *args, **kwargs):
9+
if not cls._instance:
10+
cls._instance = super(OpenAIAdapter, cls).__new__(cls)
11+
return cls._instance
12+
13+
def __init__(self, api_key: str = None):
14+
if self._initialized:
15+
return
16+
17+
if api_key is None:
18+
raise ValueError("API Key is required")
19+
20+
self.api_key = api_key
21+
self.base_url = "https://openrouter.ai/api/v1"
22+
self.client = None
23+
self.is_openrouter = False
24+
self._initialized = True
25+
self._initialize_client()
26+
self.client.embeddings = EmbeddingModel(provider_type=self.get_provider_type()).embeddings
27+
28+
def __call__(self, *args, **kwargs):
29+
if self.client is None:
30+
raise RuntimeError("Client is not initialized")
31+
return self.client
32+
33+
def _initialize_client(self):
34+
try:
35+
client = openai.OpenAI(
36+
api_key=self.api_key,
37+
)
38+
client.models.list()
39+
self.client = client
40+
except openai.AuthenticationError:
41+
try:
42+
client = openai.OpenAI(
43+
api_key=self.api_key,
44+
base_url=self.base_url,
45+
)
46+
client.models.list()
47+
self.client = client
48+
self.is_openrouter = True
49+
except openai.AuthenticationError:
50+
raise openai.AuthenticationError("Invalid API key")
51+
52+
def validate(self):
53+
return self.client is not None
54+
55+
def get_provider_type(self):
56+
return "openrouter" if self.is_openrouter else "openai"
57+
58+
def get_client(self):
59+
return self.client
60+
61+
@classmethod
62+
def reset(cls):
63+
cls._instance = None
64+
cls._initialized = False
65+

0 commit comments

Comments
 (0)