chopratejas commited on
Commit
a6eefcc
·
1 Parent(s): 62d80bc

fix(learn): pass explicit model in tests to avoid API key requirement

Browse files

SessionAnalyzer() without a model calls _detect_default_model() which
raises when no API keys are set (e.g., in CI). Pass model="test-model"
in the three tests that mock _call_llm.

headroom/learn/analyzer.py CHANGED
@@ -62,9 +62,7 @@ class SessionAnalyzer:
62
  def __init__(self, model: str | None = None):
63
  self.model = model
64
 
65
- def analyze(
66
- self, project: ProjectInfo, sessions: list[SessionData]
67
- ) -> AnalysisResult:
68
  """Analyze sessions and produce recommendations via LLM."""
69
  all_calls = [tc for s in sessions for tc in s.tool_calls]
70
  failed_calls = [tc for tc in all_calls if tc.is_error]
@@ -134,7 +132,9 @@ def _build_digest(project: ProjectInfo, sessions: list[SessionData]) -> str:
134
 
135
  for session in sessions:
136
  if chars_used > char_budget:
137
- lines.append(f"... (remaining {len(sessions) - sessions.index(session)} sessions truncated)")
 
 
138
  break
139
 
140
  session_header = (
@@ -179,7 +179,7 @@ def _format_event(event: SessionEvent) -> str | None:
179
 
180
  if event.type == "user_message" and event.text.strip():
181
  text = event.text.strip()[:300]
182
- return f" [{event.msg_index}] USER: \"{text}\""
183
 
184
  if event.type == "interruption":
185
  return f" [{event.msg_index}] INTERRUPTED: {event.text[:150]}"
@@ -188,7 +188,7 @@ def _format_event(event: SessionEvent) -> str | None:
188
  return (
189
  f" [{event.msg_index}] SUBAGENT: {event.agent_tool_count} tool calls, "
190
  f"{event.agent_tokens:,} tokens, {event.agent_duration_ms / 1000:.1f}s "
191
- f"— prompt: \"{event.agent_prompt[:100]}\""
192
  )
193
 
194
  return None
@@ -384,7 +384,5 @@ class FailureAnalyzer:
384
  def __init__(self) -> None:
385
  self._analyzer = SessionAnalyzer()
386
 
387
- def analyze(
388
- self, project: ProjectInfo, sessions: list[SessionData]
389
- ) -> AnalysisResult:
390
  return self._analyzer.analyze(project, sessions)
 
62
  def __init__(self, model: str | None = None):
63
  self.model = model
64
 
65
+ def analyze(self, project: ProjectInfo, sessions: list[SessionData]) -> AnalysisResult:
 
 
66
  """Analyze sessions and produce recommendations via LLM."""
67
  all_calls = [tc for s in sessions for tc in s.tool_calls]
68
  failed_calls = [tc for tc in all_calls if tc.is_error]
 
132
 
133
  for session in sessions:
134
  if chars_used > char_budget:
135
+ lines.append(
136
+ f"... (remaining {len(sessions) - sessions.index(session)} sessions truncated)"
137
+ )
138
  break
139
 
140
  session_header = (
 
179
 
180
  if event.type == "user_message" and event.text.strip():
181
  text = event.text.strip()[:300]
182
+ return f' [{event.msg_index}] USER: "{text}"'
183
 
184
  if event.type == "interruption":
185
  return f" [{event.msg_index}] INTERRUPTED: {event.text[:150]}"
 
188
  return (
189
  f" [{event.msg_index}] SUBAGENT: {event.agent_tool_count} tool calls, "
190
  f"{event.agent_tokens:,} tokens, {event.agent_duration_ms / 1000:.1f}s "
191
+ f'— prompt: "{event.agent_prompt[:100]}"'
192
  )
193
 
194
  return None
 
384
  def __init__(self) -> None:
385
  self._analyzer = SessionAnalyzer()
386
 
387
+ def analyze(self, project: ProjectInfo, sessions: list[SessionData]) -> AnalysisResult:
 
 
388
  return self._analyzer.analyze(project, sessions)
headroom/learn/scanner.py CHANGED
@@ -240,9 +240,7 @@ class ClaudeCodeScanner(ConversationScanner):
240
  total_input_tokens += usage.get("cache_creation_input_tokens", 0)
241
  total_output_tokens += usage.get("output_tokens", 0)
242
  elif line_type == "user":
243
- self._extract_tool_results(
244
- d, tool_uses, tool_calls, events, msg_index, ts
245
- )
246
  self._extract_user_events(d, events, msg_index, ts)
247
 
248
  except (OSError, UnicodeDecodeError) as e:
@@ -251,12 +249,8 @@ class ClaudeCodeScanner(ConversationScanner):
251
 
252
  # Also wrap tool_calls as events for unified access
253
  for tc in tool_calls:
254
- if not any(
255
- e.type == "tool_call" and e.tool_call is tc for e in events
256
- ):
257
- events.append(
258
- SessionEvent(type="tool_call", msg_index=tc.msg_index, tool_call=tc)
259
- )
260
  events.sort(key=lambda e: e.msg_index)
261
 
262
  return SessionData(
@@ -332,7 +326,9 @@ class ClaudeCodeScanner(ConversationScanner):
332
  )
333
  tool_calls.append(tc)
334
  events.append(
335
- SessionEvent(type="tool_call", msg_index=msg_index, timestamp=timestamp, tool_call=tc)
 
 
336
  )
337
 
338
  # Extract subagent summary from toolUseResult metadata
 
240
  total_input_tokens += usage.get("cache_creation_input_tokens", 0)
241
  total_output_tokens += usage.get("output_tokens", 0)
242
  elif line_type == "user":
243
+ self._extract_tool_results(d, tool_uses, tool_calls, events, msg_index, ts)
 
 
244
  self._extract_user_events(d, events, msg_index, ts)
245
 
246
  except (OSError, UnicodeDecodeError) as e:
 
249
 
250
  # Also wrap tool_calls as events for unified access
251
  for tc in tool_calls:
252
+ if not any(e.type == "tool_call" and e.tool_call is tc for e in events):
253
+ events.append(SessionEvent(type="tool_call", msg_index=tc.msg_index, tool_call=tc))
 
 
 
 
254
  events.sort(key=lambda e: e.msg_index)
255
 
256
  return SessionData(
 
326
  )
327
  tool_calls.append(tc)
328
  events.append(
329
+ SessionEvent(
330
+ type="tool_call", msg_index=msg_index, timestamp=timestamp, tool_call=tc
331
+ )
332
  )
333
 
334
  # Extract subagent summary from toolUseResult metadata
tests/test_learn/test_analyzer.py CHANGED
@@ -261,7 +261,7 @@ class TestSessionAnalyzer:
261
  "memory_file_rules": [],
262
  }
263
 
264
- analyzer = SessionAnalyzer()
265
  sessions = [
266
  SessionData(
267
  session_id="s1",
@@ -283,7 +283,7 @@ class TestSessionAnalyzer:
283
  def test_handles_llm_failure_gracefully(self, mock_call_llm: MagicMock):
284
  mock_call_llm.side_effect = RuntimeError("API key not set")
285
 
286
- analyzer = SessionAnalyzer()
287
  sessions = [
288
  SessionData(
289
  session_id="s1",
@@ -309,7 +309,7 @@ class TestSessionAnalyzer:
309
  ]
310
  sessions = [SessionData(session_id="s1", tool_calls=[tc], events=events)]
311
 
312
- analyzer = SessionAnalyzer()
313
  analyzer.analyze(_project(), sessions)
314
 
315
  # Check that the digest passed to the LLM includes user message
 
261
  "memory_file_rules": [],
262
  }
263
 
264
+ analyzer = SessionAnalyzer(model="test-model")
265
  sessions = [
266
  SessionData(
267
  session_id="s1",
 
283
  def test_handles_llm_failure_gracefully(self, mock_call_llm: MagicMock):
284
  mock_call_llm.side_effect = RuntimeError("API key not set")
285
 
286
+ analyzer = SessionAnalyzer(model="test-model")
287
  sessions = [
288
  SessionData(
289
  session_id="s1",
 
309
  ]
310
  sessions = [SessionData(session_id="s1", tool_calls=[tc], events=events)]
311
 
312
+ analyzer = SessionAnalyzer(model="test-model")
313
  analyzer.analyze(_project(), sessions)
314
 
315
  # Check that the digest passed to the LLM includes user message