""" 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()