Files

793 lines
28 KiB
Python
Executable File

#!/usr/bin/env python3
"""
Complete API Processing Pipeline
=================================
This script processes OpenAPI specs through the complete workflow:
1. Smart consolidate schemas (unify duplicates + error responses)
2. Generate Go client via ogen
3. Generate client_ext.go wrapper
Usage:
cd /path/to/remnawave-api-go
python3 scripts/pipeline.py specs/api-2-3-0.json
"""
import json
import subprocess
import sys
import re
from pathlib import Path
from typing import Dict, List, Tuple
from smart_consolidate import SmartConsolidator, InlineSchemaExtractor, unify_error_responses, fix_nullable_without_type
class Colors:
HEADER = '\033[95m'
BLUE = '\033[94m'
CYAN = '\033[96m'
GREEN = '\033[92m'
YELLOW = '\033[93m'
RED = '\033[91m'
END = '\033[0m'
BOLD = '\033[1m'
def print_step(step: int, total: int, title: str):
"""Print a step header"""
print(f"\n{Colors.BOLD}{Colors.CYAN}{'='*70}")
print(f"STEP {step}/{total}: {title}")
print(f"{'='*70}{Colors.END}\n")
def print_success(message: str):
print(f"{Colors.GREEN}{message}{Colors.END}")
def print_warning(message: str):
print(f"{Colors.YELLOW}{message}{Colors.END}")
def print_error(message: str):
print(f"{Colors.RED}{message}{Colors.END}")
def print_info(message: str):
print(f"{Colors.BLUE}{message}{Colors.END}")
# ============================================================================
# STEP 1: SMART CONSOLIDATE SCHEMAS
# ============================================================================
def smart_consolidate_schemas(input_file: str, output_file: str, skip_inline_extraction: bool = False) -> Tuple[int, int, dict]:
"""
Consolidate duplicate schemas using smart analysis.
Combines old Steps 1 (consolidate) and 2 (rename) into one step.
"""
print_info(f"Loading {input_file}...")
with open(input_file, 'r') as f:
spec = json.load(f)
original_count = len(spec.get('components', {}).get('schemas', {}))
print_info("Analyzing schemas with SmartConsolidator...")
consolidator = SmartConsolidator(spec)
# Analyze duplicates
report = consolidator.analyze_duplicates()
print_info(f"Found {report['exact']['count']} exact duplicate groups ({report['exact']['total_schemas']} schemas)")
print_info(f"Found {report['structural']['count']} structural duplicate groups")
if report['near_duplicates']['count'] > 0:
print_warning(f"Found {report['near_duplicates']['count']} near-duplicate groups (metadata differs)")
if report['constraint_only']['count'] > 0:
print_warning(f"Found {report['constraint_only']['count']} constraint-only groups (validation differs)")
# Consolidate
rename_map, stats = consolidator.consolidate()
if not rename_map:
print_warning("No duplicates to consolidate")
return original_count, original_count, {}
# Apply consolidation
new_spec = consolidator.apply_consolidation(rename_map)
# Unify error responses
print_info("Unifying error responses...")
new_spec, error_stats = unify_error_responses(new_spec)
if error_stats['total_replaced'] > 0:
print_info(f"Unified {error_stats['total_replaced']} error responses (400: {error_stats['responses_unified'].get('400', 0)}, 401: {error_stats['responses_unified'].get('401', 0)})")
stats['unified_errors'] = error_stats['total_replaced']
# Fix nullable properties without type (ogen requires type for nullable fields)
print_info("Fixing nullable properties without type...")
new_spec, nullable_fixed = fix_nullable_without_type(new_spec)
if nullable_fixed > 0:
print_info(f"Fixed {nullable_fixed} nullable properties without type")
stats['nullable_fixed'] = nullable_fixed
# Extract inline schemas for reuse (optional - can cause conflicts in some specs)
if not skip_inline_extraction:
print_info("Extracting inline schemas for reuse...")
extractor = InlineSchemaExtractor(new_spec)
new_spec, extract_stats = extractor.extract_inline_schemas()
if extract_stats['extracted_count'] > 0:
print_info(f"Extracted {extract_stats['extracted_count']} inline schemas")
stats['extracted_schemas'] = extract_stats['extracted_count']
else:
print_info("Skipping inline schema extraction")
print_info(f"Writing {output_file}...")
with open(output_file, 'w') as f:
json.dump(new_spec, f, indent=2, ensure_ascii=False)
# Print top consolidated groups
print_info("Top consolidated groups:")
for name, schemas in sorted(stats['consolidated_names'].items(), key=lambda x: -len(x[1]))[:5]:
print(f" {name} <- {len(schemas)} schemas")
new_count = len(new_spec.get('components', {}).get('schemas', {}))
stats['final_count'] = new_count
print_success(f"Consolidated {original_count}{new_count} schemas (-{original_count - new_count}, -{(original_count-new_count)*100//original_count}%)")
return original_count, new_count, stats
# ============================================================================
# STEP 1.5: PATCH SPEC FOR TEXT/PLAIN SUBSCRIPTION ENDPOINTS
# ============================================================================
# These subscription endpoints return text/plain (subscription configs as strings),
# but the OpenAPI spec doesn't declare response content, causing ogen to skip them.
SUBSCRIPTION_TEXT_OPERATIONS = [
'SubscriptionController_getSubscription',
'SubscriptionController_getSubscriptionByClientType',
'SubscriptionController_getSubscriptionWithType',
]
def patch_subscription_text_responses(spec: dict) -> int:
"""
Patch the spec to add text/plain response content
for subscription endpoints that return raw subscription configs.
Modifies spec in-place. Returns the number of operations patched.
"""
patched = 0
for path, path_item in spec.get('paths', {}).items():
for http_method, op in path_item.items():
if not isinstance(op, dict):
continue
op_id = op.get('operationId', '')
if op_id not in SUBSCRIPTION_TEXT_OPERATIONS:
continue
responses = op.get('responses', {})
resp_200 = responses.get('200', {})
# Add text/plain content if not already present
if 'content' not in resp_200:
resp_200['content'] = {}
if 'text/plain' not in resp_200['content']:
resp_200['content']['text/plain'] = {
'schema': {'type': 'string'}
}
patched += 1
print_info(f"Patched {op_id} with text/plain response")
responses['200'] = resp_200
op['responses'] = responses
return patched
# ============================================================================
# STEP 1.6: SHORTEN OPERATION IDS
# ============================================================================
def shorten_operation_ids(spec: dict) -> int:
"""
Strip 'Controller' from all operationIds to produce shorter Go type names.
E.g. SubscriptionController_getSubscription → Subscription_getSubscription
Modifies spec in-place. Returns the number of operations renamed.
"""
renamed = 0
for path, path_item in spec.get('paths', {}).items():
for http_method, op in path_item.items():
if not isinstance(op, dict):
continue
op_id = op.get('operationId', '')
if 'Controller' in op_id:
op['operationId'] = op_id.replace('Controller', '')
renamed += 1
return renamed
# ============================================================================
# STEP 1.7: STRIP 'Dto' SUFFIX FROM SCHEMA NAMES
# ============================================================================
def strip_dto_suffix(spec: dict) -> int:
"""
Remove 'Dto' suffix from all schema names and update all $ref pointers.
E.g. CreateUserRequestDto → CreateUserRequest
Modifies spec in-place. Returns the number of schemas renamed.
"""
schemas = spec.get('components', {}).get('schemas', {})
rename_map = {}
for name in list(schemas.keys()):
if name.endswith('Dto'):
new_name = name[:-3]
# Avoid collision with existing schema
if new_name not in schemas and new_name not in rename_map.values():
rename_map[name] = new_name
if not rename_map:
return 0
# Rename schemas
new_schemas = {}
for name, schema in schemas.items():
new_name = rename_map.get(name, name)
new_schemas[new_name] = schema
spec['components']['schemas'] = new_schemas
# Update all $ref pointers throughout the spec
old_prefix = '#/components/schemas/'
ref_map = {f'{old_prefix}{old}': f'{old_prefix}{new}' for old, new in rename_map.items()}
def _update_refs(obj):
if isinstance(obj, dict):
if '$ref' in obj and obj['$ref'] in ref_map:
obj['$ref'] = ref_map[obj['$ref']]
for v in obj.values():
_update_refs(v)
elif isinstance(obj, list):
for item in obj:
_update_refs(item)
_update_refs(spec)
return len(rename_map)
# ============================================================================
# STEP 1.8: FIX NUMERIC QUERY PARAMETERS THAT SHOULD BE INTEGERS
# ============================================================================
# Query parameter names that are semantically integers (pagination, limits, counts)
INTEGER_QUERY_PARAMS = {'size', 'start', 'topUsersLimit', 'topNodesLimit', 'limit', 'offset', 'page', 'count'}
def fix_number_query_params(spec: dict) -> int:
"""
Change query parameters with type 'number' to 'integer' when they represent
pagination or limit values. The upstream OpenAPI spec incorrectly uses 'number'
for these, which produces float64 in Go instead of int.
Modifies spec in-place. Returns the number of parameters fixed.
"""
fixed = 0
for path, path_item in spec.get('paths', {}).items():
for http_method, op in path_item.items():
if not isinstance(op, dict):
continue
for param in op.get('parameters', []):
if param.get('in') != 'query':
continue
schema = param.get('schema', {})
if schema.get('type') == 'number' and param.get('name') in INTEGER_QUERY_PARAMS:
schema['type'] = 'integer'
fixed += 1
return fixed
# ============================================================================
# STEP 2: GENERATE GO CLIENT WITH OGEN
# ============================================================================
def generate_ogen_client(spec_file: str) -> bool:
"""Generate Go client using ogen"""
print_info(f"Running ogen with {spec_file}...")
try:
result = subprocess.run(
[
'go', 'run', 'github.com/ogen-go/ogen/cmd/ogen@v1.19.0',
'--config', '.ogen.yml',
'--target', 'api',
'--package', 'api',
'--clean',
spec_file
],
capture_output=True,
text=True,
timeout=120
)
if result.returncode == 0:
print_success(f"Go client generated from {spec_file}")
return True
else:
print_error(f"ogen generation failed: {result.stderr}")
return False
except subprocess.TimeoutExpired:
print_error("ogen generation timed out")
return False
except Exception as e:
print_error(f"Error running ogen: {e}")
return False
# ============================================================================
# STEP 3: GENERATE CLIENT_EXT.GO
# ============================================================================
def parse_oas_client_methods(client_file: str) -> dict:
"""Parse method signatures from oas_client_gen.go"""
with open(client_file, 'r') as f:
content = f.read()
methods = {}
pattern = r'func \(c \*Client\) (\w+)\((ctx context\.Context(?:,\s*[^)]+)?)\)\s*\(([^)]+)\)'
for match in re.finditer(pattern, content, re.MULTILINE):
method_name = match.group(1)
if method_name in ['requestURL'] or method_name.startswith('send'):
continue
full_params = match.group(2)
returns = match.group(3)
# Parse params (skip ctx and variadic options)
params_list = []
has_options = False
if ', ' in full_params:
params_str = full_params.split(', ', 1)[1]
# Detect variadic ...RequestOption
if '...RequestOption' in params_str:
has_options = True
# Remove variadic param before parsing regular params
params_str = re.sub(r',?\s*options\s+\.\.\.RequestOption', '', params_str).strip()
for param in re.findall(r'(\w+)\s+([\*\w\.]+)', params_str):
params_list.append((param[0], param[1]))
returns_list = [r.strip() for r in returns.split(',')]
methods[method_name] = {
'params': params_list,
'returns': returns_list,
'has_options': has_options,
}
return methods
def parse_params_structs(params_file: str) -> dict:
"""Parse Params struct fields from oas_parameters_gen.go"""
with open(params_file, 'r') as f:
content = f.read()
params_structs = {}
# Match struct definitions with their fields
# Pattern: type XXXParams struct {\n\tField Type\n}
pattern = r'type (\w+Params) struct \{([^}]*)\}'
for match in re.finditer(pattern, content, re.DOTALL):
struct_name = match.group(1)
fields_block = match.group(2)
fields = []
# Parse fields: Name Type or Name Type `json:"..."`
for line in fields_block.strip().split('\n'):
line = line.strip()
if not line or line.startswith('//'):
continue
# Match field: UUID string or Size OptFloat64
field_match = re.match(r'^(\w+)\s+([\w\.\*\[\]]+)', line)
if field_match:
field_name = field_match.group(1)
field_type = field_match.group(2)
fields.append((field_name, field_type))
params_structs[struct_name] = fields
return params_structs
def simplify_param_type(param_type: str) -> str:
"""Convert ogen types to simpler Go types for method signatures"""
# OptString -> string, OptFloat64 -> float64, etc.
type_map = {
'OptString': 'string',
'OptInt': 'int',
'OptFloat64': 'float64',
'OptBool': 'bool',
}
return type_map.get(param_type, param_type)
# Go reserved keywords that cannot be used as identifiers
GO_KEYWORDS = {
'break', 'case', 'chan', 'const', 'continue', 'default', 'defer', 'else',
'fallthrough', 'for', 'func', 'go', 'goto', 'if', 'import', 'interface',
'map', 'package', 'range', 'return', 'select', 'struct', 'switch', 'type',
'var',
}
def safe_param_name(name: str) -> str:
"""Convert a field name to a safe Go parameter name, avoiding reserved keywords."""
lower = name.lower()
if lower in GO_KEYWORDS:
return lower + 'Val'
return lower
def _to_pascal(s: str) -> str:
"""Convert first letter to uppercase, preserving camelCase."""
if not s:
return s
return s[0].upper() + s[1:]
def parse_operations(spec_file: str) -> dict:
"""Parse operations from OpenAPI spec"""
with open(spec_file, 'r') as f:
spec = json.load(f)
operations_by_controller = {}
for path, path_item in spec.get('paths', {}).items():
for http_method, op_spec in path_item.items():
if http_method not in ['get', 'post', 'put', 'patch', 'delete']:
continue
op_id = op_spec.get('operationId')
if not op_id or '_' not in op_id:
continue
parts = op_id.split('_', 1)
controller_full = parts[0]
method_snake = parts[1]
controller = controller_full.replace('Controller', '')
method_parts = method_snake.split('_')
method_pascal = ''.join(_to_pascal(p) for p in method_parts)
go_method = controller_full + method_pascal
if controller not in operations_by_controller:
operations_by_controller[controller] = []
operations_by_controller[controller].append({
'operationId': op_id,
'goMethod': go_method,
'displayMethod': method_pascal
})
return operations_by_controller
def generate_client_ext(spec_file: str, client_file: str, output_file: str) -> Tuple[int, int]:
"""Generate client_ext.go wrapper with simplified method signatures"""
print_info("Parsing oas_client_gen.go...")
methods = parse_oas_client_methods(client_file)
print_success(f"Found {len(methods)} client methods")
# Parse params structs for simplification
params_file = client_file.replace('oas_client_gen.go', 'oas_parameters_gen.go')
print_info("Parsing oas_parameters_gen.go...")
params_structs = parse_params_structs(params_file)
print_success(f"Found {len(params_structs)} param structs")
print_info("Parsing operations from spec...")
operations_by_controller = parse_operations(spec_file)
total_ops = sum(len(ops) for ops in operations_by_controller.values())
print_success(f"Found {total_ops} operations in {len(operations_by_controller)} controllers")
def to_camel(s):
return s[0].lower() + s[1:] if s else s
def can_simplify_params(params_type: str) -> tuple:
"""
Check if Params struct can be simplified to individual arguments.
Returns (can_simplify, [(field_name, field_type, simple_type), ...])
"""
struct_name = params_type.lstrip('*')
if struct_name not in params_structs:
return False, []
fields = params_structs[struct_name]
if not fields:
return False, []
# Only simplify if all fields are simple types
simple_types = {'string', 'int', 'int64', 'float64', 'bool',
'OptString', 'OptInt', 'OptInt64', 'OptFloat64', 'OptBool'}
simplified = []
for field_name, field_type in fields:
if field_type in simple_types or field_type.startswith('Opt'):
simple = simplify_param_type(field_type)
simplified.append((field_name, field_type, simple))
else:
# Complex type, don't simplify
return False, []
return True, simplified
# Generate code
code = '''// Code generated by pipeline.py. DO NOT EDIT manually.
package api
import "context"
// ClientExt wraps the base Client with organized sub-client access.
// Use controller methods (e.g., client.Users().GetByUuid()) to call API operations.
type ClientExt struct {
\tclient *Client
'''
for controller in sorted(operations_by_controller.keys()):
field_name = to_camel(controller)
code += f'\t{field_name} *{controller}Client\n'
code += '''}
// NewClientExt creates a new ClientExt wrapper.
func NewClientExt(client *Client) *ClientExt {
\treturn &ClientExt{
\t\tclient: client,
'''
for controller in sorted(operations_by_controller.keys()):
field_name = to_camel(controller)
code += f'\t\t{field_name}: New{controller}Client(client),\n'
code += '''\t}
}
// Client returns the underlying ogen Client.
func (ce *ClientExt) Client() *Client {
\treturn ce.client
}
'''
for controller in sorted(operations_by_controller.keys()):
field_name = to_camel(controller)
code += f'''// {controller} returns the {controller}Client.
func (ce *ClientExt) {controller}() *{controller}Client {{
\treturn ce.{field_name}
}}
'''
matched_methods = 0
for controller in sorted(operations_by_controller.keys()):
code += f'''// {controller}Client provides {controller} operations.
type {controller}Client struct {{
\tclient *Client
}}
// New{controller}Client creates a new {controller}Client.
func New{controller}Client(client *Client) *{controller}Client {{
\treturn &{controller}Client{{client: client}}
}}
'''
for op in sorted(operations_by_controller[controller], key=lambda x: x['goMethod']):
go_method = op['goMethod']
display_method = op['displayMethod']
op_id = op['operationId']
if go_method not in methods:
continue
matched_methods += 1
method_info = methods[go_method]
params = method_info['params']
returns = method_info['returns']
has_options = method_info.get('has_options', False)
# options suffix for signature and call
opts_sig = ', options ...RequestOption' if has_options else ''
opts_call = ', options...' if has_options else ''
# Check if we can simplify Params struct to individual args
simplified_params = None
params_index = None
for i, (pname, ptype) in enumerate(params):
if ptype.endswith('Params'):
can_simplify, simplified = can_simplify_params(ptype)
if can_simplify:
simplified_params = simplified
params_index = i
break
if returns:
ret_type = ', '.join(returns)
if len(returns) > 1:
ret_type = f'({ret_type})'
else:
ret_type = ''
# Generate method with simplified params or original
if simplified_params and params_index is not None:
params_type = params[params_index][1]
sig_parts = []
for i, (pname, ptype) in enumerate(params):
if i == params_index:
for field_name, field_type, simple_type in simplified_params:
sig_parts.append(f'{safe_param_name(field_name)} {simple_type}')
else:
sig_parts.append(f'{pname} {ptype}')
simple_args = ', '.join(sig_parts)
params_init = f'{params_type}{{\n'
for field_name, field_type, simple_type in simplified_params:
arg_name = safe_param_name(field_name)
if field_type.startswith('Opt'):
params_init += f'\t\t{field_name}: NewOpt{simple_type.title()}({arg_name}),\n'
else:
params_init += f'\t\t{field_name}: {arg_name},\n'
params_init += '\t}'
call_args = []
for i, (pname, ptype) in enumerate(params):
if i == params_index:
call_args.append(params_init)
else:
call_args.append(pname)
code += f'''// {display_method} calls {op_id}.
func (sc *{controller}Client) {display_method}(ctx context.Context, {simple_args}{opts_sig}) {ret_type} {{
\treturn sc.client.{go_method}(ctx, {', '.join(call_args)}{opts_call})
}}
'''
else:
# Original params
if params:
params_sig = ', '.join([f'{p[0]} {p[1]}' for p in params])
params_call = ', '.join([p[0] for p in params])
else:
params_sig = ''
params_call = ''
code += f'''// {display_method} calls {op_id}.
func (sc *{controller}Client) {display_method}(ctx context.Context'''
if params_sig:
code += f', {params_sig}'
code += opts_sig + ')'
if ret_type:
code += f' {ret_type}'
code += ' {\n'
if returns:
code += f'\treturn sc.client.{go_method}(ctx'
else:
code += f'\tsc.client.{go_method}(ctx'
if params_call:
code += f', {params_call}'
code += opts_call + ')\n}\n\n'
print_info(f"Writing {output_file}...")
with open(output_file, 'w') as f:
f.write(code)
print_success(f"Generated {matched_methods}/{total_ops} methods")
return len(operations_by_controller), matched_methods
# ============================================================================
# MAIN PIPELINE
# ============================================================================
def main():
if len(sys.argv) < 2:
print_error("Usage: python3 pipeline.py <input_spec.json>")
sys.exit(1)
input_spec = sys.argv[1]
if not Path(input_spec).exists():
print_error(f"File not found: {input_spec}")
sys.exit(1)
print(f"{Colors.BOLD}{Colors.HEADER}")
print("="*70)
print(" API PROCESSING PIPELINE")
print("="*70)
print(f"{Colors.END}")
print(f"Input: {input_spec}")
# File paths - now we only need one output file since smart_consolidate does both steps
final_file = input_spec.replace('.json', '-final.json')
client_gen_file = 'api/oas_client_gen.go'
client_ext_file = 'api/client_ext.go'
try:
# Step 1: Smart consolidate (combines old Steps 1 & 2)
print_step(1, 3, "SMART CONSOLIDATE SCHEMAS")
orig_count, new_count, stats = smart_consolidate_schemas(input_spec, final_file)
# Step 1.5: Post-process the consolidated spec (in-memory)
print_info("Post-processing consolidated spec...")
with open(final_file, 'r') as f:
final_spec = json.load(f)
patched_count = patch_subscription_text_responses(final_spec)
if patched_count > 0:
print_success(f"Patched {patched_count} subscription endpoints with text/plain response")
renamed_count = shorten_operation_ids(final_spec)
if renamed_count > 0:
print_success(f"Shortened {renamed_count} operationIds (removed 'Controller')")
dto_count = strip_dto_suffix(final_spec)
if dto_count > 0:
print_success(f"Stripped 'Dto' suffix from {dto_count} schema names")
int_count = fix_number_query_params(final_spec)
if int_count > 0:
print_success(f"Fixed {int_count} query parameters: number → integer")
with open(final_file, 'w') as f:
json.dump(final_spec, f, indent=2, ensure_ascii=False)
# Step 2: Generate with ogen
print_step(2, 3, "GENERATE GO CLIENT WITH OGEN")
if not generate_ogen_client(final_file):
print_error("Failed to generate Go client")
sys.exit(1)
# Step 3: Generate client_ext
print_step(3, 3, "GENERATE CLIENT_EXT.GO WRAPPER")
ctrl_count, method_count = generate_client_ext(final_file, client_gen_file, client_ext_file)
# Summary
print(f"\n{Colors.BOLD}{Colors.GREEN}")
print("="*70)
print(" PIPELINE COMPLETED SUCCESSFULLY")
print("="*70)
print(f"{Colors.END}")
print(f"\n{Colors.BOLD}Results:{Colors.END}")
print(f" • Schemas: {orig_count}{new_count} (-{orig_count - new_count}, -{(orig_count-new_count)*100//orig_count}%)")
print(f" • Groups: {stats.get('duplicate_groups', 0)} consolidated")
print(f" • Controllers: {ctrl_count}")
print(f" • Methods: {method_count}")
print(f"\n{Colors.BOLD}Generated files:{Colors.END}")
print(f" • {final_file}")
print(f" • {client_gen_file}")
print(f" • {client_ext_file}")
print()
except Exception as e:
print_error(f"Pipeline failed: {e}")
import traceback
traceback.print_exc()
sys.exit(1)
if __name__ == '__main__':
main()