diff --git a/scrapegraphai/graphs/code_generator_graph.py b/scrapegraphai/graphs/code_generator_graph.py index 21cb51e07..facbf813c 100644 --- a/scrapegraphai/graphs/code_generator_graph.py +++ b/scrapegraphai/graphs/code_generator_graph.py @@ -82,6 +82,7 @@ def _create_graph(self) -> BaseGraph: output=["doc"], node_config={ "llm_model": self.llm_model, + "script_creator": True, "force": self.config.get("force", False), "cut": self.config.get("cut", True), "loader_kwargs": self.config.get("loader_kwargs", {}), @@ -121,7 +122,7 @@ def _create_graph(self) -> BaseGraph: ) html_analyzer_node = HtmlAnalyzerNode( - input="refined_prompt & original_html", + input="refined_prompt & doc", output=["html_info", "reduced_html"], node_config={ "llm_model": self.llm_model, diff --git a/scrapegraphai/nodes/generate_code_node.py b/scrapegraphai/nodes/generate_code_node.py index cd3a6a6bc..58157631f 100644 --- a/scrapegraphai/nodes/generate_code_node.py +++ b/scrapegraphai/nodes/generate_code_node.py @@ -16,6 +16,7 @@ from langchain_ollama import ChatOllama from langchain_core.output_parsers import StrOutputParser from langchain_core.prompts import PromptTemplate +from pydantic import ValidationError from ..prompts import TEMPLATE_INIT_CODE_GENERATION, TEMPLATE_SEMANTIC_COMPARISON from ..utils import ( @@ -120,7 +121,7 @@ def execute(self, state: dict) -> dict: reduced_html = input_data[3] answer = input_data[4] - self.raw_html = state["original_html"][0].page_content + self.raw_html = state["doc"][0].page_content simplefied_schema = str(transform_schema(self.output_schema.schema())) @@ -373,8 +374,20 @@ def semantic_comparison( Dict[str, Any]: A dictionary containing the comparison result, differences, and explanation. """ - reference_result_dict = self.output_schema(**reference_result).dict() - if are_content_equal(generated_result, reference_result_dict): + try: + reference_data = self.output_schema.model_validate( + reference_result + ).model_dump() + except (ValidationError, TypeError): + # JsonOutputParser can return JSON that does not match the schema. + # Keep the complete reference for the semantic comparison below. + reference_data = reference_result + + if ( + isinstance(generated_result, dict) + and isinstance(reference_data, dict) + and are_content_equal(generated_result, reference_data) + ): return { "are_semantically_equivalent": True, "differences": [], @@ -412,7 +425,7 @@ def semantic_comparison( return chain.invoke( { "generated_result": json.dumps(generated_result, indent=2), - "reference_result": json.dumps(reference_result_dict, indent=2), + "reference_result": json.dumps(reference_data, indent=2), } ) diff --git a/tests/test_code_generator_graph.py b/tests/test_code_generator_graph.py new file mode 100644 index 000000000..16b4f0591 --- /dev/null +++ b/tests/test_code_generator_graph.py @@ -0,0 +1,138 @@ +"""Offline regressions for the code generator's HTML and reference state.""" + +import json + +import pytest +from langchain_core.documents import Document +from langchain_core.messages import AIMessage +from langchain_core.outputs import ChatGeneration, ChatResult +from langchain_core.runnables import RunnableLambda +from langchain_openai import ChatOpenAI +from pydantic import BaseModel + +from scrapegraphai.graphs import CodeGeneratorGraph +from scrapegraphai.nodes import GenerateCodeNode + + +class Project(BaseModel): + title: str + description: str = "" + + +class Projects(BaseModel): + projects: list[Project] + + +HTML = '