generator / prompt_preprocessor.py
Gia Bao
Initial deploy: XAI Terraform Generator
da4f611
Raw
History Blame Contribute Delete
15.6 kB
"""
Prompt Preprocessor for Terraform Code Generation
===================================================
Uses TF-KB to enhance raw prompts before sending to LLM.
Pipeline:
1. Entity Extraction: NL → resource types (via TF-KB)
2. Knowledge Enrichment: inject required args, dependencies
3. Prompt Augmentation: format enhanced prompt
Usage:
from prompt_preprocessor import PromptPreprocessor
pp = PromptPreprocessor()
enhanced = pp.enhance_prompt("Create an S3 bucket with versioning")
"""
import re
from tf_knowledge_base import TerraformKB
class PromptPreprocessor:
def __init__(self, confidence_threshold: float = 0.65):
self.kb = TerraformKB()
self.confidence_threshold = confidence_threshold
def extract_resource_types(self, prompt: str,
ground_truth_resources: list[str] = None) -> list[str]:
"""
Extract resource types from prompt.
If ground_truth_resources provided (from IaC-Eval dataset), use those directly.
Otherwise, use NL extraction from TF-KB.
"""
if ground_truth_resources:
# Validate against KB — catch hallucinated types
validated = []
for rt in ground_truth_resources:
rt = rt.strip()
if self.kb.resource_exists(rt):
validated.append(rt)
else:
# Try to find the correct resource type
similar = self.kb.find_similar_resources(rt, top_k=1)
if similar and similar[0][1] > 0.8:
validated.append(similar[0][0])
return list(dict.fromkeys(validated)) # deduplicate, preserve order
# NL extraction
candidates = self.kb.extract_resources(prompt)
return [c['resource_type'] for c in candidates
if c['confidence'] >= self.confidence_threshold]
def build_resource_info(self, resource_types: list[str]) -> str:
"""Build structured resource info block for prompt injection."""
# Resolve dependencies (depth=1 only — direct deps, no transitive)
all_resources = set(resource_types)
for rt in resource_types:
for dep in self.kb.get_dependencies(rt):
all_resources.add(dep)
self._last_all_resources = all_resources
lines = []
idx = 1
for rt in sorted(all_resources):
schema = self.kb.get_resource_schema(rt)
if not schema:
continue
required = [a["name"] for a in schema.get("required", [])]
# Include type info for required args
required_with_types = []
for a in schema.get("required", []):
atype = a.get("type", "string")
required_with_types.append(f"{a['name']} ({atype})")
nested = [n for n in schema.get("nested_blocks", {}).keys()
if n != "timeouts"]
# Mark which nested blocks are required (min_items > 0)
required_blocks = []
optional_blocks = []
for bname, binfo in schema.get("nested_blocks", {}).items():
if bname == "timeouts":
continue
if binfo.get("min_items", 0) > 0:
required_blocks.append(bname)
else:
optional_blocks.append(bname)
is_dependency = rt not in resource_types
prefix = f"{idx}. {rt}"
if is_dependency:
prefix += " (auto-added dependency)"
lines.append(prefix)
if required:
lines.append(f" Required arguments: {', '.join(required_with_types)}")
if required_blocks:
lines.append(f" Required blocks: {', '.join(required_blocks)}")
if optional_blocks:
# Only show relevant optional blocks (max 5)
shown = optional_blocks[:5]
lines.append(f" Optional blocks: {', '.join(shown)}")
# Show nested block details (required args inside blocks + sub-blocks)
for bname, binfo in schema.get("nested_blocks", {}).items():
if bname == "timeouts":
continue
block_required = [f"{a['name']} ({a.get('type', 'string')})"
for a in binfo.get("required", [])]
block_optional = [a['name'] for a in binfo.get("optional", [])]
# Check required sub-blocks (min_items > 0)
required_sub = [sb for sb, si in binfo.get("nested_blocks", {}).items()
if si.get("min_items", 0) > 0]
if block_required or block_optional or required_sub:
detail = f" Block '{bname}':"
if block_required:
detail += f" required=[{', '.join(block_required)}]"
if required_sub:
detail += f" must contain: [{', '.join(required_sub)}]"
if block_optional:
detail += f" optional=[{', '.join(block_optional[:4])}]"
lines.append(detail)
# Show key optional args for common resources (helps LLM pick correct args)
key_optional = self._get_key_optional_args(rt, schema)
if key_optional:
lines.append(f" Key optional args: {', '.join(key_optional)}")
idx += 1
# Dependencies
dep_lines = []
for rt in sorted(all_resources):
deps = self.kb.get_dependencies(rt)
deps_in_scope = [d for d in deps if d in all_resources]
if deps_in_scope:
for dep in deps_in_scope:
# Find the linking argument
link_arg = self._find_link_arg(rt, dep)
if link_arg:
dep_lines.append(f" {rt}.{link_arg}{dep}")
else:
dep_lines.append(f" {rt} → depends on {dep}")
if dep_lines:
lines.append("")
lines.append("Resource dependencies:")
lines.extend(dep_lines)
return "\n".join(lines)
def _get_key_optional_args(self, resource_type: str, schema: dict) -> list[str]:
"""Return key optional arguments that LLMs commonly need for a resource."""
# Curated map of important optional args per resource
KEY_ARGS = {
"aws_instance": ["ami", "instance_type", "subnet_id", "vpc_security_group_ids", "tags"],
"aws_s3_bucket": ["bucket", "tags"],
"aws_security_group": ["name", "description", "ingress", "egress", "tags"],
"aws_vpc": ["cidr_block", "enable_dns_support", "enable_dns_hostnames", "tags"],
"aws_subnet": ["vpc_id", "cidr_block", "availability_zone", "map_public_ip_on_launch", "tags"],
"aws_db_instance": ["engine", "engine_version", "instance_class", "allocated_storage", "username", "password", "db_subnet_group_name", "skip_final_snapshot"],
"aws_lambda_function": ["function_name", "runtime", "handler", "role", "filename", "source_code_hash"],
"aws_iam_role": ["name", "assume_role_policy", "tags"],
"aws_iam_policy": ["name", "policy", "description"],
"aws_lb": ["name", "internal", "load_balancer_type", "subnets", "security_groups"],
"aws_lb_target_group": ["name", "port", "protocol", "vpc_id", "target_type"],
"aws_lb_listener": ["load_balancer_arn", "port", "protocol", "default_action"],
"aws_ecs_cluster": ["name"],
"aws_ecs_task_definition": ["family", "container_definitions", "network_mode", "requires_compatibilities", "cpu", "memory"],
"aws_ecs_service": ["name", "cluster", "task_definition", "desired_count", "launch_type"],
"aws_cloudwatch_log_group": ["name", "retention_in_days"],
"aws_route53_zone": ["name"],
"aws_route53_record": ["zone_id", "name", "type", "ttl", "records"],
"aws_elastic_beanstalk_application": ["name", "description"],
"aws_elastic_beanstalk_environment": ["name", "application", "solution_stack_name"],
"aws_launch_template": ["name_prefix", "image_id", "instance_type"],
"aws_autoscaling_group": ["min_size", "max_size", "desired_capacity", "vpc_zone_identifier"],
"aws_sns_topic": ["name"],
"aws_sqs_queue": ["name"],
"aws_dynamodb_table": ["name", "billing_mode", "hash_key"],
"aws_eip": ["domain"],
"aws_nat_gateway": ["allocation_id", "subnet_id"],
"aws_internet_gateway": ["vpc_id", "tags"],
"aws_route_table": ["vpc_id", "tags"],
"aws_cloudfront_distribution": ["enabled", "origin", "default_cache_behavior"],
"aws_acm_certificate": ["domain_name", "validation_method"],
}
return KEY_ARGS.get(resource_type, [])
def _find_link_arg(self, resource_type: str, dependency_type: str) -> str | None:
"""Guess the linking argument between resource and its dependency."""
# Common patterns
dep_short = dependency_type.replace("aws_", "")
schema = self.kb.get_resource_schema(resource_type)
if not schema:
return None
all_args = [a["name"] for a in schema.get("required", [])] + \
[a["name"] for a in schema.get("optional", [])]
# Look for *_id patterns
for arg in all_args:
if dep_short.endswith(arg.replace("_id", "").replace("_arn", "")):
return arg
if arg in (f"{dep_short}_id", f"{dep_short}_arn",
f"{dep_short.split('_')[-1]}_id"):
return arg
# Common specific mappings
specific = {
("aws_route53_record", "aws_route53_zone"): "zone_id",
("aws_subnet", "aws_vpc"): "vpc_id",
("aws_security_group", "aws_vpc"): "vpc_id",
("aws_internet_gateway", "aws_vpc"): "vpc_id",
("aws_db_instance", "aws_db_subnet_group"): "db_subnet_group_name",
("aws_s3_bucket_versioning", "aws_s3_bucket"): "bucket",
("aws_s3_bucket_policy", "aws_s3_bucket"): "bucket",
("aws_s3_object", "aws_s3_bucket"): "bucket",
("aws_iam_instance_profile", "aws_iam_role"): "role",
("aws_iam_role_policy_attachment", "aws_iam_role"): "role",
("aws_iam_role_policy_attachment", "aws_iam_policy"): "policy_arn",
("aws_lambda_function", "aws_iam_role"): "role",
("aws_elastic_beanstalk_environment", "aws_elastic_beanstalk_application"): "application",
("aws_lb_listener", "aws_lb"): "load_balancer_arn",
("aws_lb_target_group", "aws_vpc"): "vpc_id",
("aws_nat_gateway", "aws_subnet"): "subnet_id",
("aws_nat_gateway", "aws_eip"): "allocation_id",
("aws_route_table_association", "aws_route_table"): "route_table_id",
("aws_route_table_association", "aws_subnet"): "subnet_id",
("aws_ecs_service", "aws_ecs_cluster"): "cluster",
("aws_ecs_service", "aws_ecs_task_definition"): "task_definition",
}
return specific.get((resource_type, dependency_type))
def enhance_prompt(self, prompt: str,
ground_truth_resources: list[str] = None) -> str:
"""
Main entry point: enhance a raw prompt with TF-KB knowledge.
Args:
prompt: Raw user prompt
ground_truth_resources: Optional list of resource types from dataset
Returns:
Enhanced prompt with resource schemas and dependencies
"""
resource_types = self.extract_resource_types(prompt, ground_truth_resources)
if not resource_types:
# Can't extract any resources — return original prompt with generic guidance
return (
"Generate valid Terraform HCL code for AWS infrastructure. "
"Use only real AWS provider resource types (e.g., aws_instance, aws_s3_bucket). "
"Include all required arguments for each resource.\n\n"
f"Task: {prompt}"
)
resource_info = self.build_resource_info(resource_types)
# Add example snippets for detected resources
from tf_examples import get_examples_for_resources
examples = get_examples_for_resources(list(self._last_all_resources) if hasattr(self, '_last_all_resources') else resource_types)
# Add composition templates for resource combinations
from tf_compositions import get_composition_template
compositions = get_composition_template(list(self._last_all_resources) if hasattr(self, '_last_all_resources') else resource_types)
enhanced = (
f"Generate Terraform HCL code for the following AWS infrastructure.\n"
f"The resource schemas and examples below are for REFERENCE — use them as guidance "
f"but include any additional resources needed for a complete, working configuration.\n\n"
f"{resource_info}\n\n"
)
if compositions:
enhanced += f"IMPORTANT — Follow this composition pattern for a complete setup:\n\n{compositions}\n\n"
if examples:
enhanced += f"Individual resource reference (for syntax only):\n\n{examples}\n\n"
enhanced += (
f"User request: {prompt}\n\n"
f"Important:\n"
f"- The schemas above are reference — freely add related resources (IAM roles, security groups, gateways, etc.) as needed\n"
f"- Include all required arguments for each resource you create\n"
f"- Connect resources using proper references (e.g., aws_vpc.main.id)\n"
f"- Do NOT include provider or terraform blocks\n"
f"- Use variables (var.xxx) for sensitive values like passwords, not hardcoded strings\n"
f"- Include security best practices: encryption, security groups, least-privilege IAM\n"
f"- Output ONLY valid HCL code, no explanations"
)
return enhanced
if __name__ == "__main__":
pp = PromptPreprocessor()
test_cases = [
{
"prompt": "Create an S3 bucket with versioning enabled",
"resources": None # NL extraction
},
{
"prompt": "Create a Route53 zone and add a DNS A record pointing to an ELB",
"resources": ["aws_route53_zone", "aws_route53_record", "aws_elb"]
},
{
"prompt": "Deploy an Elastic Beanstalk application with IAM role",
"resources": ["aws_elastic_beanstalk_application", "aws_elastic_beanstalk_environment",
"aws_iam_role", "aws_iam_instance_profile", "aws_iam_role_policy_attachment"]
},
]
for tc in test_cases:
print("=" * 70)
print(f"RAW PROMPT: {tc['prompt']}")
print("-" * 70)
enhanced = pp.enhance_prompt(tc['prompt'], tc.get('resources'))
print(enhanced)
print()