mirror of
https://github.com/The-Pocket/PocketFlow-Tutorial-Codebase-Knowledge.git
synced 2026-08-30 00:50:32 +08:00
make the way to count abstractions more robust
This commit is contained in:
@@ -247,29 +247,32 @@ class AnalyzeRelationships(Node):
|
|||||||
language = shared.get("language", "english") # Get language
|
language = shared.get("language", "english") # Get language
|
||||||
use_cache = shared.get("use_cache", True) # Get use_cache flag, default to True
|
use_cache = shared.get("use_cache", True) # Get use_cache flag, default to True
|
||||||
|
|
||||||
|
# Get the actual number of abstractions directly
|
||||||
|
num_abstractions = len(abstractions)
|
||||||
|
|
||||||
# Create context with abstraction names, indices, descriptions, and relevant file snippets
|
# Create context with abstraction names, indices, descriptions, and relevant file snippets
|
||||||
context = "Identified Abstractions:\n"
|
context = "Identified Abstractions:\\n"
|
||||||
all_relevant_indices = set()
|
all_relevant_indices = set()
|
||||||
abstraction_info_for_prompt = []
|
abstraction_info_for_prompt = []
|
||||||
for i, abstr in enumerate(abstractions):
|
for i, abstr in enumerate(abstractions):
|
||||||
# Use 'files' which contains indices directly
|
# Use 'files' which contains indices directly
|
||||||
file_indices_str = ", ".join(map(str, abstr["files"]))
|
file_indices_str = ", ".join(map(str, abstr["files"]))
|
||||||
# Abstraction name and description might be translated already
|
# Abstraction name and description might be translated already
|
||||||
info_line = f"- Index {i}: {abstr['name']} (Relevant file indices: [{file_indices_str}])\n Description: {abstr['description']}"
|
info_line = f"- Index {i}: {abstr['name']} (Relevant file indices: [{file_indices_str}])\\n Description: {abstr['description']}"
|
||||||
context += info_line + "\n"
|
context += info_line + "\\n"
|
||||||
abstraction_info_for_prompt.append(
|
abstraction_info_for_prompt.append(
|
||||||
f"{i} # {abstr['name']}"
|
f"{i} # {abstr['name']}"
|
||||||
) # Use potentially translated name here too
|
) # Use potentially translated name here too
|
||||||
all_relevant_indices.update(abstr["files"])
|
all_relevant_indices.update(abstr["files"])
|
||||||
|
|
||||||
context += "\nRelevant File Snippets (Referenced by Index and Path):\n"
|
context += "\\nRelevant File Snippets (Referenced by Index and Path):\\n"
|
||||||
# Get content for relevant files using helper
|
# Get content for relevant files using helper
|
||||||
relevant_files_content_map = get_content_for_indices(
|
relevant_files_content_map = get_content_for_indices(
|
||||||
files_data, sorted(list(all_relevant_indices))
|
files_data, sorted(list(all_relevant_indices))
|
||||||
)
|
)
|
||||||
# Format file content for context
|
# Format file content for context
|
||||||
file_context_str = "\n\n".join(
|
file_context_str = "\\n\\n".join(
|
||||||
f"--- File: {idx_path} ---\n{content}"
|
f"--- File: {idx_path} ---\\n{content}"
|
||||||
for idx_path, content in relevant_files_content_map.items()
|
for idx_path, content in relevant_files_content_map.items()
|
||||||
)
|
)
|
||||||
context += file_context_str
|
context += file_context_str
|
||||||
@@ -277,15 +280,21 @@ class AnalyzeRelationships(Node):
|
|||||||
return (
|
return (
|
||||||
context,
|
context,
|
||||||
"\n".join(abstraction_info_for_prompt),
|
"\n".join(abstraction_info_for_prompt),
|
||||||
|
num_abstractions, # Pass the actual count
|
||||||
project_name,
|
project_name,
|
||||||
language,
|
language,
|
||||||
use_cache,
|
use_cache,
|
||||||
) # Return use_cache
|
) # Return use_cache
|
||||||
|
|
||||||
def exec(self, prep_res):
|
def exec(self, prep_res):
|
||||||
context, abstraction_listing, project_name, language, use_cache = (
|
(
|
||||||
prep_res # Unpack use_cache
|
context,
|
||||||
)
|
abstraction_listing,
|
||||||
|
num_abstractions, # Receive the actual count
|
||||||
|
project_name,
|
||||||
|
language,
|
||||||
|
use_cache,
|
||||||
|
) = prep_res # Unpack use_cache
|
||||||
print(f"Analyzing relationships using LLM...")
|
print(f"Analyzing relationships using LLM...")
|
||||||
|
|
||||||
# Add language instruction and hints only if not English
|
# Add language instruction and hints only if not English
|
||||||
@@ -335,7 +344,7 @@ relationships:
|
|||||||
|
|
||||||
Now, provide the YAML output:
|
Now, provide the YAML output:
|
||||||
"""
|
"""
|
||||||
response = call_llm(prompt)
|
response = call_llm(prompt, use_cache=use_cache)
|
||||||
|
|
||||||
# --- Validation ---
|
# --- Validation ---
|
||||||
yaml_str = response.strip().split("```yaml")[1].split("```")[0].strip()
|
yaml_str = response.strip().split("```yaml")[1].split("```")[0].strip()
|
||||||
@@ -354,7 +363,6 @@ Now, provide the YAML output:
|
|||||||
|
|
||||||
# Validate relationships structure
|
# Validate relationships structure
|
||||||
validated_relationships = []
|
validated_relationships = []
|
||||||
num_abstractions = len(abstraction_listing.split("\n"))
|
|
||||||
for rel in relationships_data["relationships"]:
|
for rel in relationships_data["relationships"]:
|
||||||
# Check for 'label' key
|
# Check for 'label' key
|
||||||
if not isinstance(rel, dict) or not all(
|
if not isinstance(rel, dict) or not all(
|
||||||
|
|||||||
+3
-3
@@ -47,7 +47,7 @@ def call_llm(prompt: str, use_cache: bool = True) -> str:
|
|||||||
logger.info(f"RESPONSE: {cache[prompt]}")
|
logger.info(f"RESPONSE: {cache[prompt]}")
|
||||||
return cache[prompt]
|
return cache[prompt]
|
||||||
|
|
||||||
# Call the LLM if not in cache or cache disabled
|
# # Call the LLM if not in cache or cache disabled
|
||||||
# client = genai.Client(
|
# client = genai.Client(
|
||||||
# vertexai=True,
|
# vertexai=True,
|
||||||
# # TODO: change to your own project id and location
|
# # TODO: change to your own project id and location
|
||||||
@@ -59,8 +59,8 @@ def call_llm(prompt: str, use_cache: bool = True) -> str:
|
|||||||
client = genai.Client(
|
client = genai.Client(
|
||||||
api_key=os.getenv("GEMINI_API_KEY", ""),
|
api_key=os.getenv("GEMINI_API_KEY", ""),
|
||||||
)
|
)
|
||||||
# model = os.getenv("GEMINI_MODEL", "gemini-2.5-pro-exp-03-25")
|
model = os.getenv("GEMINI_MODEL", "gemini-2.5-pro-exp-03-25")
|
||||||
model = os.getenv("GEMINI_MODEL", "gemini-2.0-flash-exp")
|
# model = os.getenv("GEMINI_MODEL", "gemini-2.5-flash-preview-04-17")
|
||||||
|
|
||||||
response = client.models.generate_content(model=model, contents=[prompt])
|
response = client.models.generate_content(model=model, contents=[prompt])
|
||||||
response_text = response.text
|
response_text = response.text
|
||||||
|
|||||||
Reference in New Issue
Block a user