File size: 20,231 Bytes
da4f611
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
"""
Terraform AWS Knowledge Base (TF-KB)
=====================================
Build & query a knowledge base from the AWS Terraform provider schema.
Used for prompt preprocessing β€” maps NL prompts to exact resource types,
required arguments, and dependency chains.

Usage:
    from tf_knowledge_base import TerraformKB
    kb = TerraformKB()
    
    # Find resources mentioned in a prompt
    resources = kb.extract_resources("Create an S3 bucket with versioning")
    # β†’ [{'resource_type': 'aws_s3_bucket', 'confidence': 0.9, ...}]
    
    # Get schema for a resource
    schema = kb.get_resource_schema('aws_s3_bucket')
    # β†’ {'required': [...], 'optional': [...], 'nested_blocks': {...}}
    
    # Get dependencies
    deps = kb.get_dependencies('aws_route53_record')
    # β†’ ['aws_route53_zone']
"""

import json
import os
import re
from pathlib import Path
from difflib import SequenceMatcher

KB_DIR = Path(__file__).parent / "tf_kb"
REGISTRY_PATH = KB_DIR / "resource_registry.json"
NL_MAPPING_PATH = KB_DIR / "nl_mapping.json"
DEP_GRAPH_PATH = KB_DIR / "dependency_graph.json"

# Load dependency graph from auto-generated JSON
# core_dependencies: always inject (IAM roles, parent resources)
# optional_dependencies: only inject when prompt mentions related features
def _load_dependency_map():
    if DEP_GRAPH_PATH.exists():
        with open(DEP_GRAPH_PATH) as f:
            data = json.load(f)
        # Prefer core_dependencies (filtered), fall back to full dependencies
        return data.get("core_dependencies", data.get("dependencies", {}))
    # Fallback to minimal hardcoded map if file missing
    return {
    "aws_route53_record": ["aws_route53_zone"],
    "aws_route53_zone_association": ["aws_route53_zone", "aws_vpc"],
    "aws_route53_vpc_association_authorization": ["aws_route53_zone", "aws_vpc"],
    "aws_route53_query_log": ["aws_route53_zone", "aws_cloudwatch_log_group"],
    "aws_subnet": ["aws_vpc"],
    "aws_security_group": ["aws_vpc"],
    "aws_security_group_rule": ["aws_security_group"],
    "aws_internet_gateway": ["aws_vpc"],
    "aws_nat_gateway": ["aws_subnet", "aws_eip"],
    "aws_route_table": ["aws_vpc"],
    "aws_route_table_association": ["aws_route_table", "aws_subnet"],
    "aws_route": ["aws_route_table"],
    "aws_network_acl": ["aws_vpc"],
    "aws_instance": ["aws_subnet", "aws_security_group"],
    "aws_lb": ["aws_subnet", "aws_security_group"],
    "aws_lb_target_group": ["aws_vpc"],
    "aws_lb_listener": ["aws_lb"],
    "aws_lb_listener_rule": ["aws_lb_listener"],
    "aws_db_instance": ["aws_db_subnet_group"],
    "aws_db_subnet_group": ["aws_subnet"],
    "aws_elasticache_cluster": ["aws_elasticache_subnet_group"],
    "aws_elasticache_subnet_group": ["aws_subnet"],
    "aws_elastic_beanstalk_environment": ["aws_elastic_beanstalk_application", "aws_iam_instance_profile"],
    "aws_elastic_beanstalk_application_version": ["aws_elastic_beanstalk_application", "aws_s3_bucket", "aws_s3_object"],
    "aws_iam_instance_profile": ["aws_iam_role"],
    "aws_iam_role_policy_attachment": ["aws_iam_role", "aws_iam_policy"],
    "aws_iam_role_policy": ["aws_iam_role"],
    "aws_iam_policy_document": [],
    "aws_lambda_function": ["aws_iam_role"],
    "aws_lambda_permission": ["aws_lambda_function"],
    "aws_lambda_event_source_mapping": ["aws_lambda_function"],
    "aws_ecs_service": ["aws_ecs_cluster", "aws_ecs_task_definition"],
    "aws_ecs_task_definition": ["aws_iam_role"],
    "aws_eks_node_group": ["aws_eks_cluster", "aws_subnet", "aws_iam_role"],
    "aws_s3_bucket_versioning": ["aws_s3_bucket"],
    "aws_s3_bucket_server_side_encryption_configuration": ["aws_s3_bucket"],
    "aws_s3_bucket_lifecycle_configuration": ["aws_s3_bucket"],
    "aws_s3_bucket_logging": ["aws_s3_bucket"],
    "aws_s3_bucket_policy": ["aws_s3_bucket"],
    "aws_s3_bucket_public_access_block": ["aws_s3_bucket"],
    "aws_s3_bucket_acl": ["aws_s3_bucket"],
    "aws_s3_bucket_cors_configuration": ["aws_s3_bucket"],
    "aws_s3_bucket_website_configuration": ["aws_s3_bucket"],
    "aws_s3_bucket_notification": ["aws_s3_bucket"],
    "aws_s3_object": ["aws_s3_bucket"],
    "aws_cloudwatch_log_resource_policy": ["aws_cloudwatch_log_group"],
    "aws_cloudwatch_metric_alarm": [],
    "aws_cloudfront_distribution": [],
    "aws_acm_certificate_validation": ["aws_acm_certificate", "aws_route53_record"],
    "aws_autoscaling_group": ["aws_launch_template"],
    "aws_launch_template": ["aws_security_group"],
    "aws_codebuild_project": ["aws_iam_role"],
    "aws_codepipeline": ["aws_iam_role", "aws_s3_bucket"],
    "aws_sfn_state_machine": ["aws_iam_role"],
    "aws_api_gateway_rest_api": [],
    "aws_api_gateway_resource": ["aws_api_gateway_rest_api"],
    "aws_api_gateway_method": ["aws_api_gateway_rest_api", "aws_api_gateway_resource"],
    "aws_kms_alias": ["aws_kms_key"],
    }

DEPENDENCY_MAP = _load_dependency_map()


class TerraformKB:
    def __init__(self):
        with open(REGISTRY_PATH) as f:
            data = json.load(f)
        self.resources = data["resources"]
        self.data_sources = data["data_sources"]
        self.resource_types = set(self.resources.keys())
        
        with open(NL_MAPPING_PATH) as f:
            self.nl_mapping = json.load(f)
    
    def get_resource_schema(self, resource_type: str) -> dict | None:
        """Get full schema for a resource type."""
        return self.resources.get(resource_type)
    
    def get_required_args(self, resource_type: str) -> list[str]:
        """Get required argument names for a resource."""
        schema = self.resources.get(resource_type, {})
        args = [a["name"] for a in schema.get("required", [])]
        # Also check nested blocks with min_items > 0
        for bname, binfo in schema.get("nested_blocks", {}).items():
            if binfo.get("min_items", 0) > 0:
                args.append(f"{bname} {{...}}")
        return args
    
    def get_optional_args(self, resource_type: str) -> list[str]:
        """Get optional argument names for a resource."""
        schema = self.resources.get(resource_type, {})
        return [a["name"] for a in schema.get("optional", [])]
    
    def get_nested_blocks(self, resource_type: str) -> list[str]:
        """Get nested block names for a resource."""
        schema = self.resources.get(resource_type, {})
        return list(schema.get("nested_blocks", {}).keys())
    
    def get_dependencies(self, resource_type: str) -> list[str]:
        """Get common dependencies for a resource type."""
        return DEPENDENCY_MAP.get(resource_type, [])
    
    def resource_exists(self, resource_type: str) -> bool:
        """Check if a resource type exists in the AWS provider."""
        return resource_type in self.resource_types
    
    def find_similar_resources(self, name: str, top_k: int = 5) -> list[tuple[str, float]]:
        """Find similar resource types by fuzzy matching."""
        scores = []
        for rtype in self.resource_types:
            score = SequenceMatcher(None, name.lower(), rtype.lower()).ratio()
            scores.append((rtype, score))
        scores.sort(key=lambda x: x[1], reverse=True)
        return scores[:top_k]
    
    def extract_resources(self, prompt: str) -> list[dict]:
        """
        Extract likely AWS resource types from a natural language prompt.
        Returns list of {'resource_type': str, 'confidence': float, 'matched_text': str}
        """
        prompt_lower = prompt.lower()
        candidates = {}  # resource_type β†’ best match info
        
        # Strategy 1: Direct resource type mention (aws_xxx_yyy)
        direct_matches = re.findall(r'aws_[a-z_]+', prompt_lower)
        for match in direct_matches:
            if match in self.resource_types:
                candidates[match] = {'confidence': 1.0, 'matched_text': match, 'strategy': 'direct'}
        
        # Strategy 2: NL keyword matching
        # Check bigrams first (more specific), then unigrams
        words = re.findall(r'[a-z][a-z0-9]+', prompt_lower)
        
        # Bigrams β€” prefer exact base resource over sub-resources
        for i in range(len(words) - 1):
            bigram = f"{words[i]}_{words[i+1]}"
            if bigram in self.nl_mapping:
                matched = self.nl_mapping[bigram]
                # Prefer the shortest/base resource (e.g., aws_s3_bucket over aws_s3_bucket_acl)
                base_resource = f"aws_{bigram}"
                for rtype in matched:
                    if rtype == base_resource:
                        # Exact base match gets high confidence
                        candidates[rtype] = {
                            'confidence': 0.9,
                            'matched_text': f"{words[i]} {words[i+1]}",
                            'strategy': 'bigram_exact'
                        }
                    elif rtype not in candidates:
                        # Sub-resources get lower confidence β€” only added if
                        # prompt specifically mentions them
                        candidates[rtype] = {
                            'confidence': 0.3,
                            'matched_text': f"{words[i]} {words[i+1]}",
                            'strategy': 'bigram_sub'
                        }
            # Also check as single key (e.g., "security group", "load balancer")
            phrase = f"{words[i]} {words[i+1]}"
            if phrase in self.nl_mapping:
                for rtype in self.nl_mapping[phrase]:
                    if rtype not in candidates or candidates[rtype]['confidence'] < 0.85:
                        candidates[rtype] = {
                            'confidence': 0.85,
                            'matched_text': phrase,
                            'strategy': 'phrase'
                        }
        
        # Unigrams (lower confidence, skip common English words)
        STOPWORDS = {'create', 'set', 'up', 'add', 'with', 'and', 'the', 'for',
                     'this', 'that', 'from', 'into', 'use', 'using', 'make',
                     'has', 'have', 'get', 'put', 'new', 'two', 'three', 'one',
                     'all', 'each', 'any', 'both', 'should', 'must', 'will',
                     'can', 'may', 'also', 'deploy', 'provision', 'configure',
                     'enable', 'enabled', 'disable', 'disabled', 'include',
                     'including', 'specify', 'specified', 'ensure', 'allow',
                     'name', 'named', 'type', 'value', 'key', 'default',
                     'index', 'global', 'local', 'primary', 'secondary',
                     'public', 'private', 'internal', 'external'}
        for word in words:
            # Try exact match, then singular form (strip trailing 's'/'es')
            lookup_words = [word]
            if word.endswith('ies'):
                lookup_words.append(word[:-3] + 'y')  # policies β†’ policy
            elif word.endswith('ses') or word.endswith('xes') or word.endswith('zes'):
                lookup_words.append(word[:-2])  # instances β†’ instance
            elif word.endswith('s') and not word.endswith('ss'):
                lookup_words.append(word[:-1])  # subnets β†’ subnet
            
            for lw in lookup_words:
                if lw in self.nl_mapping and len(lw) > 2 and lw not in STOPWORDS:
                    for rtype in self.nl_mapping[lw]:
                        if rtype not in candidates:
                            candidates[rtype] = {
                                'confidence': 0.5,
                                'matched_text': word,
                                'strategy': 'unigram'
                            }
                    break  # found match, don't try other forms
        
        # Strategy 3: Common service name patterns
        service_patterns = {
            r'\bs3\b': ['aws_s3_bucket'],
            r'\bec2\b': ['aws_instance'],
            r'\brds\b': ['aws_db_instance'],
            r'\bvpc\b': ['aws_vpc'],
            r'\blambda\b': ['aws_lambda_function'],
            r'\bdynamodb\b': ['aws_dynamodb_table'],
            r'\bsqs\b': ['aws_sqs_queue'],
            r'\bsns\b': ['aws_sns_topic'],
            r'\becs\b': ['aws_ecs_cluster'],
            r'\beks\b': ['aws_eks_cluster'],
            r'\becr\b': ['aws_ecr_repository'],
            r'\bkms\b': ['aws_kms_key'],
            r'\bacm\b': ['aws_acm_certificate'],
            r'\biam\b': ['aws_iam_role'],
            r'\bwaf\b': ['aws_wafv2_web_acl'],
            r'\bssm\b': ['aws_ssm_parameter'],
            r'\bsubnets?\b': ['aws_subnet'],
            r'\bsecurity.?groups?\b': ['aws_security_group'],
            r'\bnat.?gateways?\b': ['aws_nat_gateway'],
            r'\binternet.?gateways?\b': ['aws_internet_gateway'],
            r'\broute.?tables?\b': ['aws_route_table'],
            r'\belastic.?beanstalk\b': ['aws_elastic_beanstalk_application', 'aws_elastic_beanstalk_environment'],
            r'\bcloud.?front\b': ['aws_cloudfront_distribution'],
            r'\bauto.?scaling\b': ['aws_autoscaling_group'],
            r'\belastic.?cache\b': ['aws_elasticache_cluster'],
            r'\bredshift\b': ['aws_redshift_cluster'],
            r'\bcloud.?watch\b': ['aws_cloudwatch_metric_alarm'],
            r'\bcode.?pipeline\b': ['aws_codepipeline'],
            r'\bcode.?build\b': ['aws_codebuild_project'],
            r'\bstep.?functions?\b': ['aws_sfn_state_machine'],
            r'\bapi.?gateway\b': ['aws_api_gateway_rest_api'],
            r'\bfargate\b': ['aws_ecs_cluster', 'aws_ecs_service', 'aws_ecs_task_definition'],
            r'\balb\b|\bapplication.?load.?balancer\b': ['aws_lb', 'aws_lb_target_group', 'aws_lb_listener'],
            r'\bnlb\b|\bnetwork.?load.?balancer\b': ['aws_lb'],
            r'\broute.?53\b': ['aws_route53_zone'],
        }
        for pattern, rtypes in service_patterns.items():
            if re.search(pattern, prompt_lower):
                for rtype in rtypes:
                    if rtype not in candidates or candidates[rtype]['confidence'] < 0.7:
                        candidates[rtype] = {
                            'confidence': 0.7,
                            'matched_text': re.search(pattern, prompt_lower).group(),
                            'strategy': 'service_pattern'
                        }
        
        # Build result list sorted by confidence
        results = []
        for rtype, info in candidates.items():
            results.append({
                'resource_type': rtype,
                'confidence': info['confidence'],
                'matched_text': info['matched_text'],
                'strategy': info['strategy']
            })
        results.sort(key=lambda x: x['confidence'], reverse=True)
        # Filter: only return high-confidence matches by default
        # Keep sub-resources only if their base is also matched
        base_types = {r['resource_type'] for r in results if r['confidence'] >= 0.7}
        filtered = []
        for r in results:
            if r['confidence'] >= 0.5:
                filtered.append(r)
            elif r['strategy'] == 'bigram_sub':
                # Keep sub-resource only if base resource is in scope
                # e.g., keep aws_s3_bucket_versioning if aws_s3_bucket matched
                base = '_'.join(r['resource_type'].split('_')[:3])  # aws_s3_bucket
                if base in base_types:
                    continue  # skip sub-resources, they'll be suggested via nested_blocks
        return filtered
    
    def get_resource_summary(self, resource_type: str) -> str:
        """Get a concise text summary of a resource for prompt injection."""
        if resource_type not in self.resources:
            similar = self.find_similar_resources(resource_type, top_k=3)
            return f"WARNING: '{resource_type}' does not exist. Did you mean: {', '.join(s[0] for s in similar)}?"
        
        schema = self.resources[resource_type]
        required = [a["name"] for a in schema.get("required", [])]
        nested = list(schema.get("nested_blocks", {}).keys())
        # Filter out 'timeouts' from nested β€” it's always there and not useful
        nested = [n for n in nested if n != "timeouts"]
        
        deps = self.get_dependencies(resource_type)
        
        lines = [f"  {resource_type}:"]
        if required:
            lines.append(f"    Required args: {', '.join(required)}")
        if nested:
            lines.append(f"    Nested blocks: {', '.join(nested)}")
        if deps:
            lines.append(f"    Typically needs: {', '.join(deps)}")
        
        return "\n".join(lines)
    
    def build_enrichment(self, resource_types: list[str]) -> str:
        """
        Build enrichment text for a list of resource types.
        Includes schemas + auto-resolved dependencies.
        """
        all_resources = set(resource_types)
        # Auto-resolve dependencies
        to_check = list(resource_types)
        while to_check:
            rt = to_check.pop()
            for dep in self.get_dependencies(rt):
                if dep not in all_resources:
                    all_resources.add(dep)
                    to_check.append(dep)
        
        lines = ["Resources and their schemas:"]
        for rt in sorted(all_resources):
            lines.append(self.get_resource_summary(rt))
        
        # Add dependency info
        dep_lines = []
        for rt in sorted(all_resources):
            deps = self.get_dependencies(rt)
            deps_in_scope = [d for d in deps if d in all_resources]
            if deps_in_scope:
                dep_lines.append(f"  {rt} β†’ depends on: {', '.join(deps_in_scope)}")
        
        if dep_lines:
            lines.append("\nDependency chain:")
            lines.extend(dep_lines)
        
        return "\n".join(lines)


def build_kb():
    """Rebuild the knowledge base from terraform provider schema."""
    import subprocess
    
    tmpdir = "/tmp/tf-schema"
    os.makedirs(tmpdir, exist_ok=True)
    
    main_tf = os.path.join(tmpdir, "main.tf")
    with open(main_tf, "w") as f:
        f.write('''terraform {
  required_providers {
    aws = { source = "hashicorp/aws", version = "~> 5.0" }
  }
}
provider "aws" {
  region                      = "us-east-1"
  skip_credentials_validation = true
  skip_metadata_api_check     = true
  skip_requesting_account_id  = true
  access_key                  = "mock"
  secret_key                  = "mock"
}
''')
    
    print("Running terraform init...")
    subprocess.run(["terraform", "init", "-input=false"], cwd=tmpdir, capture_output=True)
    
    print("Extracting provider schema...")
    result = subprocess.run(
        ["terraform", "providers", "schema", "-json"],
        cwd=tmpdir, capture_output=True, text=True
    )
    schema = json.loads(result.stdout)
    
    # Parse and save (same logic as build script)
    provider_key = [k for k in schema['provider_schemas'] if 'aws' in k][0]
    ps = schema['provider_schemas'][provider_key]
    # ... (parsing logic)
    print(f"Schema extracted from {provider_key}")


# Quick test
if __name__ == "__main__":
    kb = TerraformKB()
    
    print(f"=== TF-KB Stats ===")
    print(f"Resources: {len(kb.resources)}")
    print(f"Data sources: {len(kb.data_sources)}")
    print()
    
    # Test extract
    test_prompts = [
        "Create an S3 bucket with versioning enabled",
        "Set up a VPC with two subnets and an internet gateway",
        "Create a Route53 zone and add a DNS record",
        "Deploy an Elastic Beanstalk application with an IAM role",
        "Create a DynamoDB table with a global secondary index",
    ]
    
    for prompt in test_prompts:
        print(f"Prompt: {prompt}")
        resources = kb.extract_resources(prompt)
        for r in resources[:5]:
            print(f"  β†’ {r['resource_type']} (conf={r['confidence']:.1f}, via {r['strategy']}: '{r['matched_text']}')")
        
        # Show enrichment for top resources
        top_types = [r['resource_type'] for r in resources[:3]]
        if top_types:
            enrichment = kb.build_enrichment(top_types)
            print(f"\n  Enrichment:\n{enrichment}")
        print()