Feat: Finished stage 3, LangChain and LangGraph foundations

This commit is contained in:
2025-12-06 16:30:12 -05:00
parent 57ef82f191
commit 480d77a18e
4 changed files with 338 additions and 32 deletions
@@ -17,3 +17,19 @@ Feature: LangChain chat provider coverage
When the provider generates changes and the graph raises an exception
Then the provider response should capture the graph failure
And the progress callback should end at 100 percent even on failure
@coverage @langchain @streaming
Scenario: LangChain provider streams node events for progress updates
Given a LangChain chat provider is configured with a fake LangChain graph
When the provider streams changes with incremental graph events
Then the provider response should contain the streamed change payload
And the progress callback should include streaming milestones
And the LangChain graph stream should receive the thread-aware configuration
And the provider should report the estimated token usage
@coverage @langchain @retry
Scenario: LangChain provider unwraps nested retry failures
Given a LangChain chat provider is configured with a fake LangChain graph
When the provider generates changes and retries exhaust with nested errors
Then the provider response should capture the nested retry failure
And the progress callback should end at 100 percent even on failure
+131 -7
View File
@@ -18,21 +18,29 @@ def step_configure_langchain_provider(context):
context.progress_callback = progress_callback
context.requested_models = []
context.llm_instance = object()
context.llm_instance = MagicMock()
context.llm_instance.get_num_tokens = MagicMock(return_value=0)
def fake_llm_factory(model_id: str):
context.requested_models.append(model_id)
return context.llm_instance
context.provider = LangChainChatProvider(
name="test-langchain-provider",
model_id="test-model",
llm_factory=fake_llm_factory,
max_retries=2,
)
def build_provider(*, supports_streaming: bool) -> LangChainChatProvider:
return LangChainChatProvider(
name="test-langchain-provider",
model_id="test-model",
llm_factory=fake_llm_factory,
max_retries=2,
supports_streaming=supports_streaming,
)
context.create_provider = build_provider
context.provider = build_provider(supports_streaming=False)
context.project = MagicMock(spec=Project)
context.plan = MagicMock(spec=Plan)
context.plan.prompt = "Implement feature"
context.contexts = [MagicMock(spec=Context)]
context.contexts[0].content = "Initial context"
@when("the provider generates changes with a validation failure response")
@@ -91,6 +99,81 @@ def step_provider_generates_exception(context):
context.failure_message = failure_message
@when("the provider streams changes with incremental graph events")
def step_provider_streams_changes(context):
context.provider = context.create_provider(supports_streaming=True)
generated_change = Change(
plan_id=7,
file_path="src/streamed.py",
operation=OperationType.CREATE,
new_content="print('streaming')",
)
validation_message = "Looks great"
streaming_events = [
{"load_context": {}},
{"analyze_requirements": {}},
{"generate_plan": {"generated_changes": [generated_change]}},
{
"validate": {
"validation_result": {"status": "PASS", "message": validation_message}
}
},
]
context.llm_instance.get_num_tokens.return_value = 321
with patch(
"cleveragents.providers.llm.langchain_chat_provider.PlanGenerationGraph"
) as graph_cls:
mock_graph = graph_cls.return_value
mock_graph.stream.return_value = iter(streaming_events)
context.graph_instance = mock_graph
context.response = context.provider.generate_changes(
context.project,
context.plan,
context.contexts,
progress_callback=context.progress_callback,
)
constructor_args, constructor_kwargs = graph_cls.call_args
stream_call_args = mock_graph.stream.call_args
context.graph_constructor_args = constructor_args
context.graph_constructor_kwargs = constructor_kwargs
context.stream_call_args = stream_call_args
context.expected_change = generated_change
context.expected_token_count = 321
@when("the provider generates changes and retries exhaust with nested errors")
def step_provider_generates_retry_failure(context):
final_message = "All retry attempts failed"
class RetryError(Exception):
def __init__(self):
super().__init__("retry failed")
self.last_attempt = _FakeAttempt()
class _FakeAttempt:
def exception(self):
return RuntimeError(final_message)
retry_error = RetryError()
with patch(
"cleveragents.providers.llm.langchain_chat_provider.PlanGenerationGraph"
) as graph_cls:
mock_graph = graph_cls.return_value
mock_graph.invoke.side_effect = retry_error
context.graph_instance = mock_graph
context.response = context.provider.generate_changes(
context.project,
context.plan,
context.contexts,
progress_callback=context.progress_callback,
)
context.failure_message = final_message
@then(
"the provider response should contain the generated change data and validation error"
)
@@ -127,6 +210,40 @@ def step_assert_graph_invocation(context):
assert context.graph_constructor_kwargs["max_retries"] == 2
@then("the provider response should contain the streamed change payload")
def step_assert_streaming_response(context):
assert context.response is not None
assert context.response.changes, "Expected streaming change results"
streamed_change = context.response.changes[0]
assert streamed_change.file_path == context.expected_change.file_path
assert streamed_change.operation == context.expected_change.operation
assert streamed_change.new_content == context.expected_change.new_content
assert context.response.error_message in (None, "")
@then("the progress callback should include streaming milestones")
def step_assert_streaming_progress(context):
assert context.progress_updates == [5, 15, 40, 70, 90, 100]
@then("the LangChain graph stream should receive the thread-aware configuration")
def step_assert_stream_call(context):
call_args, call_kwargs = context.stream_call_args
assert call_args == (
context.project,
context.plan,
context.contexts,
)
assert "thread_id" in call_kwargs
assert call_kwargs["thread_id"].startswith("provider-")
@then("the provider should report the estimated token usage")
def step_assert_token_count(context):
assert context.response is not None
assert context.response.token_count == context.expected_token_count
@then("the provider response should capture the graph failure")
def step_assert_graph_failure_response(context):
assert context.response is not None
@@ -136,6 +253,13 @@ def step_assert_graph_failure_response(context):
assert context.response.token_count == 0
@then("the provider response should capture the nested retry failure")
def step_assert_nested_retry_failure(context):
assert context.response is not None
assert context.response.changes == []
assert context.response.error_message == context.failure_message
@then("the progress callback should end at 100 percent even on failure")
def step_assert_failure_progress_updates(context):
assert context.progress_updates == [5, 100]
+4 -3
View File
@@ -4028,7 +4028,7 @@ If you can do all of the above by end of Day 1, you're on track!
- [X] Created helper script `robot/test_context_analysis.py` for complex Python logic
- [X] All tests passing (nox -s unit_tests, nox -s integration_tests)
- [X] Code: Fix LangGraph checkpointing with thread_id in config
- [ ] Code: Integrate agents into services
- [X] Code: Integrate agents into services
- [X] Update `ContextService` to use `ContextAnalysisAgent` (`src/cleveragents/application/services/context_service.py:408-515`).
- [X] Add streaming support to context commands via `analyze_context_streaming*` methods (`src/cleveragents/application/services/context_service.py:517-596`).
- [X] Add LangSmith metadata for context analysis (`src/cleveragents/application/services/context_service.py:421-633`).
@@ -4103,7 +4103,7 @@ If you can do all of the above by end of Day 1, you're on track!
- [X] Add docstrings to all agent classes (already present)
- [X] Document state TypedDict fields
- [X] Show example configurations as inline code
- [ ] Prepare for Docusaurus API reference generation (deferred to Phase 7)
- [X] Prepare for Docusaurus API reference generation (tracking moved to Phase 7 Docusaurus automation task)
- [X] Stage 2.7 Completion Criteria (COMPLETE 2025-11-30)
- [X] All Behave tests pass for plan_generation_agent_coverage.feature (29 scenarios, 225 steps - PASSING)
- [X] All Behave tests pass for plan_generation_uncovered_lines.feature (15 scenarios, 91 steps - PASSING)
@@ -4118,7 +4118,7 @@ If you can do all of the above by end of Day 1, you're on track!
- [X] 90%+ test coverage for agents package (95% overall coverage, exceeds requirement)
- [X] EntityMemory integration complete with 23 memory service scenarios passing
- [X] Stage 2.7.1 Test Alignment complete - all agent tests passing
- [ ] Stage 3: LangChain/ LangGraph foundations
- [X] Stage 3: LangChain/ LangGraph foundations
- [X] Install LangChain/LangGraph dependencies
- [X] Added to pyproject.toml under `[project.optional-dependencies.llm]`
- [X] Verified installation with `pip install -e .[llm]`
@@ -4940,6 +4940,7 @@ If you can do all of the above by end of Day 1, you're on track!
- [ ] Create reference pages for LangGraph workflow patterns in docs/docs/api/agents.md
- [ ] Generate module documentation for core packages in docs/docs/api/
- [ ] Add docstrings to all public APIs if missing
- [ ] Plan Docusaurus API reference automation (migrated from Stage 2.7.6) including script requirements, dependency mapping, and CI integration
- [ ] Code: Create architecture diagrams in docs/docs/architecture/
- [ ] Generate Mermaid diagrams for runtime architecture in docs/docs/architecture/runtime.md
- [ ] Create deployment topology diagrams in docs/docs/architecture/deployment.md
@@ -4,7 +4,9 @@ from __future__ import annotations
import uuid
from collections.abc import Callable
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, ClassVar, SupportsInt, cast
from tenacity import Retrying, stop_after_attempt, wait_exponential
from cleveragents.agents.plan_generation import PlanGenerationGraph
from cleveragents.domain.models.core import Change, Context, Plan, Project
@@ -13,13 +15,29 @@ from cleveragents.domain.providers.ai_provider import (
ProviderResponse,
)
if TYPE_CHECKING:
if TYPE_CHECKING: # pragma: no cover - imported for type checking only
from langchain_core.language_models import BaseLanguageModel
TokenEstimatorResult = int | float | SupportsInt | str | None
TokenEstimator = Callable[[str], TokenEstimatorResult]
def _is_retry_error(error: Exception) -> bool:
"""Return True if exception appears to be a Tenacity RetryError."""
return error.__class__.__name__ == "RetryError"
class LangChainChatProvider(AIProviderInterface):
"""AI provider that uses a LangChain chat model with PlanGenerationGraph."""
_DEFAULT_PROGRESS_MAP: ClassVar[dict[str, int]] = {
"load_context": 15,
"analyze_requirements": 40,
"generate_plan": 70,
"validate": 90,
}
def __init__(
self,
*,
@@ -27,11 +45,15 @@ class LangChainChatProvider(AIProviderInterface):
model_id: str,
llm_factory: Callable[[str], BaseLanguageModel],
max_retries: int = 3,
supports_streaming: bool = True,
progress_map: dict[str, int] | None = None,
) -> None:
self._name = name
self._model_id = model_id
self._llm_factory = llm_factory
self._max_retries = max_retries
self._max_retries = max(1, max_retries)
self._supports_streaming = supports_streaming
self._progress_map = progress_map or self._DEFAULT_PROGRESS_MAP.copy()
@property
def name(self) -> str: # pragma: no cover - simple accessor
@@ -50,41 +72,184 @@ class LangChainChatProvider(AIProviderInterface):
) -> ProviderResponse:
"""Generate code changes by running the LangGraph workflow."""
if progress_callback:
progress_callback(5)
llm = self._llm_factory(self._model_id)
graph = PlanGenerationGraph(llm=llm, max_retries=self._max_retries)
thread_id = f"provider-{uuid.uuid4()}"
token_count = self._estimate_token_usage(llm, plan, contexts)
state: dict[str, Any]
try:
state = graph.invoke(project, plan, contexts, thread_id=thread_id)
if progress_callback:
progress_callback(5)
if progress_callback and self._supports_streaming:
state = self._execute_with_streaming(
graph,
project,
plan,
contexts,
thread_id,
progress_callback,
)
else:
state = self._invoke_with_retry(
graph,
project,
plan,
contexts,
thread_id,
)
if progress_callback:
progress_callback(90)
progress_callback(100)
except Exception as exc: # pragma: no cover - defensive path
if progress_callback:
progress_callback(100)
error_message = self._extract_retry_error_message(exc)
return ProviderResponse(
changes=[],
model_used=self._model_id,
token_count=0,
error_message=str(exc),
token_count=token_count,
error_message=error_message or str(exc),
)
if progress_callback:
progress_callback(90)
validation = self._safe_validation_result(state)
error_message = self._safe_error_message(state)
if not error_message and validation:
status = str(validation.get("status", "")).upper()
if status == "FAIL":
error_message = validation.get("message", "Validation failed")
generated_changes: list[Change] = state.get("generated_changes", [])
validation = state.get("validation_result", {})
error = state.get("error")
if validation.get("status") == "FAIL" and not error:
error = validation.get("message", "Validation failed")
if progress_callback:
progress_callback(100)
generated_changes = self._safe_generated_changes(state)
return ProviderResponse(
changes=generated_changes,
model_used=self._model_id,
token_count=0,
error_message=error,
token_count=token_count,
error_message=error_message,
)
def _invoke_with_retry(
self,
graph: PlanGenerationGraph,
project: Project,
plan: Plan,
contexts: list[Context],
thread_id: str,
) -> dict[str, Any]:
retryer = Retrying(
stop=stop_after_attempt(self._max_retries),
wait=wait_exponential(multiplier=0.25, min=0.25, max=2.0),
reraise=True,
)
for attempt in retryer:
with attempt:
return graph.invoke(project, plan, contexts, thread_id=thread_id)
raise RuntimeError("Retries exhausted while invoking plan generation graph")
def _execute_with_streaming(
self,
graph: PlanGenerationGraph,
project: Project,
plan: Plan,
contexts: list[Context],
thread_id: str,
progress_callback: Callable[[int], None],
) -> dict[str, Any]:
state: dict[str, Any] = {
"generated_changes": [],
"validation_result": {},
"error": None,
}
try:
for event in graph.stream(project, plan, contexts, thread_id=thread_id):
node_name, payload = self._extract_event(event)
self._emit_progress(node_name, progress_callback)
if isinstance(payload, dict):
state.update(cast(dict[str, Any], payload))
progress_callback(100)
except Exception:
progress_callback(100)
raise
return state
def _extract_event(self, event: dict[str, Any]) -> tuple[str, Any]:
try:
key, payload = next(iter(event.items()))
except StopIteration: # pragma: no cover - defensive guard
return "__unknown__", {}
return key, payload
def _emit_progress(
self, node_name: str, progress_callback: Callable[[int], None]
) -> None:
percent = self._progress_map.get(node_name)
if percent is not None:
progress_callback(percent)
def _estimate_token_usage(
self,
llm: BaseLanguageModel,
plan: Plan,
contexts: list[Context],
) -> int:
get_tokens = getattr(llm, "get_num_tokens", None)
if not callable(get_tokens):
return 0
estimator = cast(TokenEstimator, get_tokens)
prompt_text = plan.prompt or ""
for ctx in contexts:
content = getattr(ctx, "content", None) or ""
if content:
prompt_text += f"\n{content[:2000]}"
try:
tokens = estimator(prompt_text)
except Exception: # pragma: no cover - defensive fallback
return 0
normalized_tokens = tokens
if normalized_tokens is None:
return 0
try:
return int(normalized_tokens)
except (TypeError, ValueError): # pragma: no cover - defensive fallback
return 0
def _safe_generated_changes(self, state: dict[str, Any]) -> list[Change]:
value = state.get("generated_changes")
if isinstance(value, list):
return cast(list[Change], value)
return []
def _safe_validation_result(self, state: dict[str, Any]) -> dict[str, Any]:
value = state.get("validation_result")
if isinstance(value, dict):
return cast(dict[str, Any], value)
return {}
def _safe_error_message(self, state: dict[str, Any]) -> str | None:
value = state.get("error")
if value is None:
return None
return str(value)
def _extract_retry_error_message(self, error: Exception) -> str | None:
if not _is_retry_error(error):
return None
last_attempt = getattr(error, "last_attempt", None)
if last_attempt is None:
return str(error)
exception_callable = getattr(last_attempt, "exception", None)
if not callable(exception_callable):
return str(error)
last_exception = exception_callable()
if isinstance(last_exception, Exception):
return str(last_exception)
return str(error)