chopratejas commited on
Commit
126b60e
·
1 Parent(s): a451bf7

Fix LangChain tool_call argument handling for varied message formats

Browse files

LangChain provides tool_call args in different shapes (dict args, str
arguments, nested function.arguments) depending on the source. Add
_tool_call_args_to_json() helper to normalize all formats to JSON strings.
Use .get() instead of [] to handle missing keys gracefully.

headroom/integrations/langchain/chat_model.py CHANGED
@@ -79,6 +79,22 @@ def _check_langchain_available() -> None:
79
  )
80
 
81
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
82
  def langchain_available() -> bool:
83
  """Check if LangChain is installed."""
84
  return LANGCHAIN_AVAILABLE
@@ -241,11 +257,11 @@ class HeadroomChatModel(BaseChatModel):
241
  if msg.tool_calls:
242
  entry["tool_calls"] = [
243
  {
244
- "id": tc["id"],
245
  "type": "function",
246
  "function": {
247
- "name": tc["name"],
248
- "arguments": json.dumps(tc["args"]),
249
  },
250
  }
251
  for tc in msg.tool_calls
@@ -928,11 +944,11 @@ def optimize_messages(
928
  if hasattr(msg, "tool_calls") and msg.tool_calls:
929
  entry["tool_calls"] = [
930
  {
931
- "id": tc["id"],
932
  "type": "function",
933
  "function": {
934
- "name": tc["name"],
935
- "arguments": json.dumps(tc["args"]),
936
  },
937
  }
938
  for tc in msg.tool_calls
 
79
  )
80
 
81
 
82
+ def _tool_call_args_to_json(tc: dict[str, Any]) -> str:
83
+ """Normalize tool call arguments to JSON string for OpenAI format.
84
+
85
+ LangChain can provide 'args' (dict) or 'arguments' (str) depending on source.
86
+ """
87
+ if "args" in tc:
88
+ val = tc["args"]
89
+ return json.dumps(val) if isinstance(val, dict) else str(val)
90
+ if "arguments" in tc:
91
+ val = tc["arguments"]
92
+ return val if isinstance(val, str) else json.dumps(val)
93
+ if "function" in tc and isinstance(tc["function"], dict):
94
+ return str(tc["function"].get("arguments", "{}"))
95
+ return "{}"
96
+
97
+
98
  def langchain_available() -> bool:
99
  """Check if LangChain is installed."""
100
  return LANGCHAIN_AVAILABLE
 
257
  if msg.tool_calls:
258
  entry["tool_calls"] = [
259
  {
260
+ "id": tc.get("id", ""),
261
  "type": "function",
262
  "function": {
263
+ "name": tc.get("name", ""),
264
+ "arguments": _tool_call_args_to_json(tc),
265
  },
266
  }
267
  for tc in msg.tool_calls
 
944
  if hasattr(msg, "tool_calls") and msg.tool_calls:
945
  entry["tool_calls"] = [
946
  {
947
+ "id": tc.get("id", ""),
948
  "type": "function",
949
  "function": {
950
+ "name": tc.get("name", ""),
951
+ "arguments": _tool_call_args_to_json(tc),
952
  },
953
  }
954
  for tc in msg.tool_calls
tests/test_integrations/langchain/test_langchain_live.py ADDED
@@ -0,0 +1,312 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Live LangChain integration tests — no mocks, real API keys from .env.
2
+
3
+ Run with:
4
+ pytest tests/test_integrations/langchain/test_langchain_live.py -v -s
5
+ # Or with env loaded:
6
+ set -a && source .env && set +a && pytest tests/test_integrations/langchain/test_langchain_live.py -v -s
7
+
8
+ Requires: OPENAI_API_KEY and/or ANTHROPIC_API_KEY in environment (e.g. from .env).
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import os
14
+ from pathlib import Path
15
+
16
+ import pytest
17
+
18
+ # Load .env from project root if present
19
+ _project_root = Path(__file__).resolve().parents[3]
20
+ _env = _project_root / ".env"
21
+ if _env.exists():
22
+ try:
23
+ from dotenv import load_dotenv
24
+
25
+ load_dotenv(_env)
26
+ except ImportError:
27
+ pass
28
+
29
+ try:
30
+ from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
31
+ from langchain_core.tools import tool
32
+
33
+ LANGCHAIN_AVAILABLE = True
34
+ except ImportError:
35
+ LANGCHAIN_AVAILABLE = False
36
+
37
+ OPENAI_KEY = os.environ.get("OPENAI_API_KEY", "").strip()
38
+ ANTHROPIC_KEY = os.environ.get("ANTHROPIC_API_KEY", "").strip()
39
+ HAS_OPENAI = bool(OPENAI_KEY)
40
+ HAS_ANTHROPIC = bool(ANTHROPIC_KEY)
41
+ HAS_ANY_KEY = HAS_OPENAI or HAS_ANTHROPIC
42
+
43
+ pytestmark = [
44
+ pytest.mark.skipif(not LANGCHAIN_AVAILABLE, reason="LangChain not installed"),
45
+ pytest.mark.skipif(
46
+ not HAS_ANY_KEY, reason="No OPENAI_API_KEY or ANTHROPIC_API_KEY in env (e.g. .env)"
47
+ ),
48
+ ]
49
+
50
+
51
+ @pytest.fixture
52
+ def openai_llm():
53
+ """Real ChatOpenAI if OPENAI_API_KEY is set."""
54
+ if not HAS_OPENAI:
55
+ pytest.skip("OPENAI_API_KEY not set")
56
+ from langchain_openai import ChatOpenAI
57
+
58
+ return ChatOpenAI(model="gpt-4o-mini", temperature=0)
59
+
60
+
61
+ @pytest.fixture
62
+ def anthropic_llm():
63
+ """Real ChatAnthropic if ANTHROPIC_API_KEY is set."""
64
+ if not HAS_ANTHROPIC:
65
+ pytest.skip("ANTHROPIC_API_KEY not set")
66
+ from langchain_anthropic import ChatAnthropic
67
+
68
+ # Allow override via env (e.g. claude-sonnet-4-20250514); default to a common current model
69
+ model = os.environ.get("ANTHROPIC_MODEL", "claude-sonnet-4-20250514")
70
+ return ChatAnthropic(model=model, temperature=0)
71
+
72
+
73
+ # --- HeadroomChatModel: invoke (sync) ---
74
+
75
+
76
+ class TestHeadroomChatModelLiveOpenAI:
77
+ """Live tests: HeadroomChatModel wrapping ChatOpenAI."""
78
+
79
+ def test_wrap_openai_and_invoke(self, openai_llm):
80
+ from headroom.integrations import HeadroomChatModel
81
+
82
+ model = HeadroomChatModel(openai_llm)
83
+ messages = [HumanMessage(content="Reply with exactly: OK")]
84
+ response = model.invoke(messages)
85
+
86
+ assert response is not None
87
+ assert hasattr(response, "content")
88
+ assert response.content is not None
89
+ assert len(response.content) > 0
90
+ assert len(model._metrics_history) >= 1
91
+ m = model._metrics_history[-1]
92
+ assert m.tokens_before >= 0
93
+ assert m.tokens_after >= 0
94
+
95
+ def test_invoke_with_string_input(self, openai_llm):
96
+ """LangChain allows invoke(str); BaseChatModel converts to messages."""
97
+ from headroom.integrations import HeadroomChatModel
98
+
99
+ model = HeadroomChatModel(openai_llm)
100
+ response = model.invoke("Say hello in one word.")
101
+ assert response is not None
102
+ assert hasattr(response, "content")
103
+ assert len(response.content) > 0
104
+
105
+ def test_system_and_user_messages(self, openai_llm):
106
+ from headroom.integrations import HeadroomChatModel
107
+
108
+ model = HeadroomChatModel(openai_llm)
109
+ messages = [
110
+ SystemMessage(content="You are a helpful assistant. Be very brief."),
111
+ HumanMessage(content="What is 2+2? One number only."),
112
+ ]
113
+ response = model.invoke(messages)
114
+ assert response.content is not None
115
+ assert "4" in response.content or "four" in response.content.lower()
116
+
117
+ def test_get_savings_summary_after_calls(self, openai_llm):
118
+ from headroom.integrations import HeadroomChatModel
119
+
120
+ model = HeadroomChatModel(openai_llm)
121
+ model.invoke([HumanMessage(content="Hi")])
122
+ summary = model.get_savings_summary()
123
+ assert summary["total_requests"] >= 1
124
+ assert "total_tokens_saved" in summary
125
+ assert "average_savings_percent" in summary
126
+
127
+
128
+ class TestHeadroomChatModelLiveAnthropic:
129
+ """Live tests: HeadroomChatModel wrapping ChatAnthropic.
130
+
131
+ If your Anthropic account does not have access to the default model,
132
+ set ANTHROPIC_MODEL=your-model (e.g. claude-3-5-sonnet-20241022) in .env.
133
+ """
134
+
135
+ def test_wrap_anthropic_and_invoke(self, anthropic_llm):
136
+ from headroom.integrations import HeadroomChatModel
137
+
138
+ model = HeadroomChatModel(anthropic_llm)
139
+ messages = [HumanMessage(content="Reply with exactly: OK")]
140
+ try:
141
+ response = model.invoke(messages)
142
+ except Exception as e:
143
+ if "404" in str(e) or "not_found" in str(e).lower():
144
+ pytest.skip(f"Anthropic model not available: {e}")
145
+ raise
146
+ assert response is not None
147
+ assert response.content is not None
148
+ assert len(response.content) > 0
149
+ assert len(model._metrics_history) >= 1
150
+
151
+ def test_provider_detection_anthropic(self, anthropic_llm):
152
+ from headroom.integrations import HeadroomChatModel
153
+
154
+ model = HeadroomChatModel(anthropic_llm)
155
+ _ = model.pipeline
156
+ assert model._provider is not None
157
+ assert "anthropic" in model._provider.__class__.__name__.lower() or "anthropic" in str(
158
+ type(model._provider)
159
+ )
160
+
161
+
162
+ # --- Streaming ---
163
+
164
+
165
+ class TestHeadroomChatModelStreamingLive:
166
+ """Live streaming tests."""
167
+
168
+ def test_stream_openai(self, openai_llm):
169
+ from headroom.integrations import HeadroomChatModel
170
+
171
+ model = HeadroomChatModel(openai_llm)
172
+ messages = [HumanMessage(content="Count from 1 to 3, one number per line.")]
173
+ chunks = list(model.stream(messages))
174
+ assert len(chunks) >= 1
175
+ full = "".join(c.content for c in chunks if c.content)
176
+ assert "1" in full or "2" in full or "3" in full
177
+
178
+ @pytest.mark.asyncio
179
+ async def test_astream_openai(self, openai_llm):
180
+ from headroom.integrations import HeadroomChatModel
181
+
182
+ model = HeadroomChatModel(openai_llm)
183
+ messages = [HumanMessage(content="Say 'stream' and nothing else.")]
184
+ count = 0
185
+ async for chunk in model.astream(messages):
186
+ if chunk.content:
187
+ count += 1
188
+ assert count >= 1
189
+
190
+
191
+ # --- Tool calling (real round-trip) ---
192
+
193
+
194
+ class TestHeadroomChatModelToolCallsLive:
195
+ """Live tool-calling tests: bind_tools + invoke with tool use."""
196
+
197
+ def test_bind_tools_and_invoke_with_tool_output(self, openai_llm):
198
+ """Simulate agent turn: user -> model (tool call) -> tool result -> model. We compress tool result."""
199
+ from headroom.integrations import HeadroomChatModel
200
+
201
+ @tool
202
+ def big_search(query: str) -> str:
203
+ """Search (returns large JSON)."""
204
+ import json
205
+
206
+ return json.dumps(
207
+ {
208
+ "results": [
209
+ {"id": i, "title": f"Result {i}", "snippet": "x" * 200} for i in range(50)
210
+ ],
211
+ "total": 50,
212
+ }
213
+ )
214
+
215
+ base = openai_llm.bind_tools([big_search])
216
+ model = HeadroomChatModel(base)
217
+
218
+ # User asks something that may trigger tool use
219
+ messages = [
220
+ HumanMessage(
221
+ content="Search for 'python tutorials' and tell me how many results you got."
222
+ ),
223
+ ]
224
+ response = model.invoke(messages)
225
+
226
+ assert response is not None
227
+ # Either direct answer or tool_calls
228
+ if response.tool_calls:
229
+ assert len(response.tool_calls) >= 1
230
+ tc = response.tool_calls[0]
231
+ assert "name" in tc or hasattr(tc, "get")
232
+ assert len(model._metrics_history) >= 1
233
+
234
+ def test_messages_with_tool_result_compressed(self, openai_llm):
235
+ """Conversation with tool call + large tool result; Headroom should compress the tool result."""
236
+ import json
237
+
238
+ from headroom.integrations import HeadroomChatModel
239
+
240
+ model = HeadroomChatModel(openai_llm)
241
+ # Simulate: user -> assistant (tool call) -> tool (large result) -> user (follow-up)
242
+ large_result = json.dumps([{"id": i, "data": "x" * 100} for i in range(100)])
243
+ messages = [
244
+ HumanMessage(content="Get items 1 to 100."),
245
+ AIMessage(
246
+ content="",
247
+ tool_calls=[
248
+ {
249
+ "id": "call_1",
250
+ "name": "get_items",
251
+ "args": {"limit": 100},
252
+ "type": "tool_call",
253
+ }
254
+ ],
255
+ ),
256
+ ToolMessage(content=large_result, tool_call_id="call_1"),
257
+ HumanMessage(content="How many items did you get? One number only."),
258
+ ]
259
+ response = model.invoke(messages)
260
+
261
+ assert response is not None
262
+ assert response.content is not None
263
+ # Optimization should have run (tool content was large)
264
+ assert len(model._metrics_history) >= 1
265
+ last = model._metrics_history[-1]
266
+ assert last.tokens_before >= last.tokens_after or last.tokens_before == last.tokens_after
267
+
268
+
269
+ # --- LCEL chain ---
270
+
271
+
272
+ class TestHeadroomLCELive:
273
+ """Live LCEL chain tests."""
274
+
275
+ def test_prompt_pipe_headroom_pipe_llm(self, openai_llm):
276
+ from langchain_core.output_parsers import StrOutputParser
277
+ from langchain_core.prompts import ChatPromptTemplate
278
+
279
+ from headroom.integrations import HeadroomChatModel
280
+
281
+ model = HeadroomChatModel(openai_llm)
282
+ prompt = ChatPromptTemplate.from_messages(
283
+ [
284
+ ("system", "You are helpful. Reply in one short sentence."),
285
+ ("human", "{input}"),
286
+ ]
287
+ )
288
+ chain = prompt | model | StrOutputParser()
289
+ result = chain.invoke({"input": "What is the capital of France?"})
290
+ assert result is not None
291
+ assert "Paris" in result or "paris" in result.lower()
292
+
293
+
294
+ # --- optimize_messages standalone (no LLM call) ---
295
+
296
+
297
+ class TestOptimizeMessagesLive:
298
+ """Live optimize_messages with real Headroom pipeline (no API key needed for this)."""
299
+
300
+ def test_optimize_messages_large_conversation(self):
301
+ from headroom.integrations import optimize_messages
302
+
303
+ messages = [SystemMessage(content="You are helpful.")]
304
+ for i in range(30):
305
+ messages.append(HumanMessage(content=f"Question {i}: What is {i}?"))
306
+ messages.append(AIMessage(content=f"Answer: {i}."))
307
+ messages.append(HumanMessage(content="Summarize the last answer."))
308
+
309
+ optimized, metrics = optimize_messages(messages)
310
+ assert len(optimized) >= 1
311
+ assert metrics["tokens_before"] >= metrics["tokens_after"]
312
+ assert "transforms_applied" in metrics