"""SQL层:WHERE约束解析 + DB输入行生成""" import re import logging import itertools logger = logging.getLogger(__name__) # ── String literal protection ── def _protect_strings(text: str) -> (str, list): """Replace string literals with placeholders. Returns (clean_text, replacements).""" replacements = [] def _repl(m): idx = len(replacements) replacements.append(m.group(0)) return f"__STR{idx}__" cleaned = re.sub(r"'[^']*'|\"[^\"]*\"", _repl, text) return cleaned, replacements def _restore_strings(text: str, replacements: list) -> str: for i, s in enumerate(replacements): text = text.replace(f"__STR{i}__", s) return text # ── Bracket-aware AND splitting ── _RE_BETWEEN_SPAN = re.compile(r'\bBETWEEN\b', re.IGNORECASE) def _scan_between_spans(text: str) -> (str, list): """Replace `X BETWEEN lo AND hi` spans with `__BTW{i}__` placeholders. The inner AND that belongs to a BETWEEN clause must not be treated as a top-level AND separator. Returns (protected_text, spans) where each span is: {subject, lo, hi, neg} — each operand is the raw token string. Subject may be a host variable (`:NAME`), a column, or a literal; lo/hi likewise (host var, column, or literal). The full span text (subject through hi) is replaced by a single placeholder token. """ spans = [] out = [] pos = 0 i = 0 while i < len(text): m = _RE_BETWEEN_SPAN.search(text, i) if not m: break # subject: token (optionally ':'-prefixed) immediately before an # optional NOT that precedes BETWEEN j = m.start() - 1 while j >= 0 and text[j].isspace(): j -= 1 before = text[max(0, j - 5):m.start()] mnot = re.search(r'NOT\s*$', before, re.IGNORECASE) if mnot: j = m.start() - len(mnot.group(0)) - 1 while j >= 0 and text[j].isspace(): j -= 1 tok_end = j while j >= 0 and (text[j].isalnum() or text[j] in '_-'): j -= 1 subject_start = j if j >= 0 and text[j] == ':' else j + 1 subject = text[subject_start:tok_end + 1] out.append(text[pos:subject_start]) # lo: from after BETWEEN to the first top-level AND j = m.end() depth = 0 and_pos = None k = j while k < len(text): ch = text[k] if ch == '(': depth += 1 elif ch == ')': depth -= 1 elif depth == 0 and re.match(r'\bAND\b', text[k:], re.IGNORECASE): and_pos = k break k += 1 if and_pos is None: out.append(text[subject_start:m.start()]) out.append(text[m.start():]) pos = len(text) break lo = text[j:and_pos].strip() # hi: from after AND until the next top-level AND/OR or end k = and_pos + 3 hi_end = len(text) depth = 0 while k < len(text): ch = text[k] if ch == '(': depth += 1 elif ch == ')': depth -= 1 elif depth == 0 and re.match(r'\b(?:AND|OR)\b', text[k:], re.IGNORECASE): hi_end = k break k += 1 hi = text[and_pos + 3:hi_end].strip() spans.append({'subject': subject, 'lo': lo, 'hi': hi, 'neg': bool(mnot)}) out.append(f"__BTW{len(spans) - 1}__ ") pos = hi_end i = hi_end out.append(text[pos:]) return ''.join(out), spans def _split_on_AND(text: str, spans: list = None, keep_placeholders: bool = False) -> list[str]: """Split WHERE clause on AND, respecting parentheses. BETWEEN ... AND ... spans are protected first so their inner AND is not treated as a separator. When `keep_placeholders` is True (internal use), BETWEEN spans stay as `__BTW{i}__` tokens so the caller can re-parse them. """ if spans is None: text, spans = _scan_between_spans(text) parts = [] current = [] depth = 0 tokens = re.split(r'(\bAND\b|\bOR\b|[()])', text, flags=re.IGNORECASE) for token in tokens: if not token.strip(): continue if token == '(': depth += 1 current.append(token) elif token == ')': depth -= 1 current.append(token) elif token.upper() == 'AND' and depth == 0: parts.append(' '.join(current).strip()) current = [] elif token.upper() == 'OR' and depth == 0: current.append(token) # OR stays as inner condition text else: current.append(token) if current: parts.append(' '.join(current).strip()) if not keep_placeholders: parts = [_restore_between_placeholders(p, spans) for p in parts] return parts def _restore_between_placeholders(part: str, spans: list) -> str: """Replace `__BTW{i}__` tokens back with their original BETWEEN text.""" def _repl(m): idx = int(m.group(1)) if idx >= len(spans): return m.group(0) s = spans[idx] neg = 'NOT ' if s['neg'] else '' return f"{s['subject']} {neg}BETWEEN {s['lo']} AND {s['hi']}" return re.sub(r'__BTW(\d+)__', _repl, part) _BETWEEN_OP = re.compile(r'(__STR\d+__|:\w[\w-]*|[\w.-]+)') def _parse_between_operand(raw: str, replacements: list) -> dict: """Parse a BETWEEN subject/lo/hi operand into a kinded dict. Returns one of: {'kind': 'host_var', 'name': 'X'} {'kind': 'literal', 'value': '...'} {'kind': 'column', 'name': 'COL'} """ raw = raw.strip().strip('()') if raw.startswith(':'): return {'kind': 'host_var', 'name': raw[1:].upper()} m = re.match(r'__STR(\d+)__$', raw) if m and int(m.group(1)) < len(replacements): return {'kind': 'literal', 'value': replacements[int(m.group(1))].strip("'\"")} if raw.startswith('__STR') and raw.endswith('__'): idx = int(raw[5:-2]) if idx < len(replacements): return {'kind': 'literal', 'value': replacements[idx].strip("'\"")} return {'kind': 'column', 'name': raw.upper()} # ── WHERE condition parsing ── _COL_OP_PAT = re.compile( r'(:?\w[\w.-]*)\s*' # column name or host variable (with optional alias prefix) r'(=|>|<|>=|<=|<>|!=|NOT\s*=)\s*' r'(:\w[\w-]*(?::\w[\w-]*)?|__STR\d+__|[\w\d.-]+)', re.IGNORECASE ) _RE_IN_CLAUSE = re.compile( r'(\w[\w.-]*)\s+(NOT\s+)?IN\s*\((.+?)\)', re.IGNORECASE ) _RE_BETWEEN = re.compile( r'(:?\w[\w.-]*)\s+(NOT\s+)?BETWEEN\s+(.+?)\s+AND\s+(.+)', re.IGNORECASE ) _RE_LIKE = re.compile( r'(\w[\w.-]*)\s+(NOT\s+)?LIKE\s+(__STR\d+__)', re.IGNORECASE ) _RE_IS_NULL = re.compile( r'(\w[\w.-]*)\s+IS\s+(NOT\s+)?NULL', re.IGNORECASE ) def _parse_where_condition(part: str, replacements: list) -> dict | None: """Parse a single WHERE condition (after AND split).""" part = part.strip() if not part: return None # IS NULL m = _RE_IS_NULL.match(part) if m: col = m.group(1).upper() neg = bool(m.group(2)) return {'col': col, 'type': 'is_null', 'neg': neg, 'op': 'IS NULL' if not neg else 'IS NOT NULL'} # IN m = _RE_IN_CLAUSE.match(part) if m: col = m.group(1).upper() neg = bool(m.group(2)) vals_text = m.group(3) # Parse values from IN list vals = [] for v in re.split(r'\s*,\s*', vals_text): v = v.strip() if v.startswith('__STR') and v.endswith('__'): idx = int(v[5:-2]) vals.append(replacements[idx].strip("'\"")) elif v.startswith(':'): vals.append({'type': 'host_var', 'host_var': v[1:].upper()}) else: vals.append(v.strip()) return { 'col': col, 'type': 'in', 'neg': neg, 'op': 'NOT IN' if neg else 'IN', 'values': vals, } # BETWEEN m = _RE_BETWEEN.match(part) if m: subj = m.group(1).upper() neg = bool(m.group(2)) lo = m.group(3).strip() hi = m.group(4).strip() return { 'type': 'between', 'neg': neg, 'op': 'BETWEEN', 'subject': _parse_between_operand(subj, replacements), 'lo': _parse_between_operand(lo, replacements), 'hi': _parse_between_operand(hi, replacements), } # LIKE m = _RE_LIKE.match(part) if m: col = m.group(1).upper() neg = bool(m.group(2)) pat_ph = m.group(3) idx = int(pat_ph[5:-2]) pattern = replacements[idx].strip("'\"") if idx < len(replacements) else pat_ph return { 'col': col, 'type': 'like', 'neg': neg, 'op': 'NOT LIKE' if neg else 'LIKE', 'pattern': pattern, } # col op value m = _COL_OP_PAT.match(part) if m: col = m.group(1).upper() op = m.group(2).upper().strip() val = m.group(3).strip() # Normalize NOT = to <> if op == 'NOT =' or op == 'NOT=': op = '<>' if val.startswith(':'): host_var = val[1:].upper() if ':' in host_var: host_var = host_var.split(':')[0] return {'col': col, 'type': 'host_var', 'host_var': host_var, 'op': op, 'literal': None} elif val.startswith('__STR') and val.endswith('__'): idx = int(val[5:-2]) _quotes = "'\"" literal = replacements[idx].strip(_quotes) if idx < len(replacements) else val return {'col': col, 'type': 'literal', 'host_var': None, 'op': op, 'literal': literal} else: return {'col': col, 'type': 'literal', 'host_var': None, 'op': op, 'literal': val} return None # ── Column name → COBOL field name ── _COLUMN_MAP = {} def guess_cobol_field(col_name: str, table: str, declared_columns: dict, column_map: dict = None) -> str: """Map SQL column name to COBOL field name. Priority: 1. DECLARE TABLE PIC alias 2. column_map 3. naming conv 4. as-is """ if column_map is None: column_map = _COLUMN_MAP # 1. DECLARE TABLE explicit PIC mapping if table in declared_columns: for c in declared_columns[table]: if c['name'] == col_name and c.get('db_type') == 'PIC': return c.get('pic', col_name) # 2. User map key = f"{table}.{col_name}" if key in column_map: return column_map[key] # 3. Naming convention: CUST_ID → CUST-ID candidate = col_name.replace('_', '-') # 4. Strip table alias prefix: A.ID → ID if '.' in candidate: candidate = candidate.split('.')[1] return candidate # ── Main constraint extraction ── def sql_extract_constraints(where_clause: str, table: str, host_vars: dict[str, str], column_map: dict[str, str], declared_columns: dict) -> list[dict]: """Parse WHERE clause into constraint list.""" if not where_clause: return [] # Protect string literals cleaned, replacements = _protect_strings(where_clause) # Protect BETWEEN spans, then split on AND protected, spans = _scan_between_spans(cleaned) and_parts = _split_on_AND(protected, spans=spans, keep_placeholders=True) constraints = [] for part in and_parts: part = part.strip() if not part: continue btw = _parse_between_placeholder(part, spans, replacements) if btw is not None: cond = btw else: cond = _parse_where_condition(part, replacements) if cond: # Map column to COBOL field col_name = cond.get('col') or (cond.get('subject', {}) or {}).get('name') if col_name and cond['type'] == 'between': cond['cobol_field'] = guess_cobol_field(col_name, table, declared_columns, column_map) elif cond.get('col'): cobol_field = guess_cobol_field(cond['col'], table, declared_columns, column_map) cond['cobol_field'] = cobol_field constraints.append(cond) else: logger.warning(f"Unparseable WHERE condition: {_restore_strings(part, replacements)}") return constraints def _parse_between_placeholder(part: str, spans: list, replacements: list) -> dict | None: """Parse a BETWEEN span from its `__BTW{i}__` placeholder. The part may contain trailing content after the placeholder (e.g. wrapped in parentheses); only the first placeholder token is interpreted. Returns a between constraint dict, or None when the part has no placeholder. """ m = re.search(r'__BTW(\d+)__', part) if not m: return None idx = int(m.group(1)) if idx >= len(spans): return None s = spans[idx] subject = _parse_between_operand(s['subject'], replacements) return { 'type': 'between', 'neg': s['neg'], 'op': 'BETWEEN', 'subject': subject, 'lo': _parse_between_operand(s['lo'], replacements), 'hi': _parse_between_operand(s['hi'], replacements), } # ── DB input row generation ── _COLUMN_DEFAULTS = { 'CHAR': lambda size: ' ' * (size or 1), 'VARCHAR': lambda size: ' ' * (size or 1), 'INTEGER': lambda _: '000000000', 'SMALLINT': lambda _: '0000', 'DECIMAL': lambda _: '000000', 'DATE': lambda _: '20260603', 'PIC': lambda _: '?', } def _format_db_value(col_info: dict, raw_val: str) -> str: db_type = col_info.get('db_type', 'CHAR') if db_type in ('CHAR', 'VARCHAR'): size = col_info.get('size', 1) default = ' ' * size elif db_type in ('INTEGER', 'SMALLINT', 'DECIMAL'): default = _COLUMN_DEFAULTS.get(db_type, lambda _: '?')(0) else: default = _COLUMN_DEFAULTS.get(db_type, lambda _: str(raw_val)[:10])(0) if raw_val is None: return default if db_type in ('INTEGER', 'SMALLINT', 'DECIMAL'): try: return str(int(raw_val)).zfill(len(default)) except ValueError: return default return str(raw_val).ljust(len(default))[:len(default)] def _make_key_unique(key_val: str, path_index: int, seen_keys: set) -> str: stripped = key_val.strip() if len(stripped) >= 5: return key_val unique = f"{path_index:03d}{stripped[:5]}" while unique in seen_keys: unique = f"{path_index:03d}{hash(key_val) % 100000:05d}" seen_keys.add(unique) return unique def collect_sql_meta(assignments: dict, declared_columns: dict, column_map: dict = None) -> list[dict]: """Collect SQL metadata from assignments. Returns list of SQL info dicts.""" sql_meta = [] seen = set() for tgt, asgn_list in assignments.items(): if isinstance(asgn_list, dict): asgn_list = [asgn_list] for asgn in asgn_list: atype = asgn.get('type', '') if not atype.startswith('exec_sql_'): continue if atype == 'exec_sql_fetch': continue key = asgn.get('sql_text', '') if key in seen: continue seen.add(key) where = asgn.get('where', '') table = asgn.get('table', '') table = _norm_table(table) where_constraints = sql_extract_constraints( where, table, {}, column_map or {}, declared_columns ) meta = dict(asgn) meta['table'] = table meta['where_constraints'] = where_constraints sql_meta.append(meta) # Order by source position so downstream SELECTs come last (reliable # runtime query order). Entries without pos keep their insertion order. sql_meta.sort(key=lambda m: m.get('pos') or 10 ** 9) return sql_meta def _path_has_sql_ok(path_cons: list) -> bool: """Check if a path requires SQLCODE = 0 (SQL succeeded).""" sql_ok = True # default: no SQLCODE constraint, assume success for pc in path_cons: if len(pc) >= 4 and pc[0] == 'SQLCODE': if pc[1] == '<>' and pc[3]: sql_ok = False if pc[1] == '=' and not pc[3]: sql_ok = False if pc[1] == '>' and pc[3]: sql_ok = False # SQLCODE = (want True) means SQL must fail if pc[1] == '=' and pc[3] and str(pc[2]).strip() not in ('0', "'0'"): sql_ok = False break return sql_ok def _infer_columns_from_where(where_cons: list) -> list[dict]: """Infer column definitions from WHERE constraints when DECLARE TABLE is missing.""" seen = {} for wc in where_cons: col_name = wc.get('col', '').split('.')[-1] if col_name and col_name not in seen: seen[col_name] = {'name': col_name, 'db_type': 'CHAR', 'size': 10} return list(seen.values()) def _parse_select_columns(select_list: str) -> list[str]: """Parse SELECT column list into individual column names. Handles: 'COL1, COL2' → ['COL1', 'COL2'] 'COL1, COL2 AS alias' → ['COL1', 'COL2'] """ if not select_list: return [] cols = [] for part in select_list.split(','): part = part.strip() # Strip table alias prefix (T.COL → COL) if '.' in part: part = part.split('.')[1] if '.' in part else part # Strip AS alias m = re.search(r'^(\w[\w.-]*)', part) if m: cols.append(m.group(1).upper()) return cols def _rec_get(rec: dict, key: str, default=''): """rec が flat / {working_storage: ...} 両方の形式をサポートする。""" if key in rec: return str(rec[key]) ws = rec.get('working_storage', {}) if isinstance(ws, dict) and key in ws: return str(ws[key]) inp = rec.get('input', {}) if isinstance(inp, dict) and key in inp: return str(inp[key]) return default def _rec_has(rec: dict, key: str) -> bool: return key in rec or key in rec.get('working_storage', {}) or key in rec.get('input', {}) def _norm_col(name: str) -> str: """Normalize a SQL column name for comparisons ('-' and '_' are equivalent).""" return str(name).upper().replace('-', '_') def _norm_table(table: str) -> str: """Normalize a SQL table name: strip DB2 schema qualifier (SCHEMA.TABLE → TABLE). core.py captures qualified names faithfully (e.g. 'SALARYDB.EMP-MASTER') so the schema part is not lost; here we reduce to the last segment so seeds target the YAML schema table (EMP-MASTER / EMP_MASTER). No-op for plain names. """ return str(table).rsplit('.', 1)[-1] def _declared_cols_for(declared_columns: dict, table: str) -> list[dict]: """Look up declared columns for a SQL table, tolerating '-'/'_' naming. YAML schema registers tables with underscores (EMP_MASTER) while SQL references may use hyphens (EMP-MASTER); missing this fallback makes build_db_input fall back to 10-char inferred sizes, padding stored keys with trailing spaces so runtime '=' lookups with unpadded host vars fail. """ cols = declared_columns.get(table, []) if not cols and '-' in table: cols = declared_columns.get(table.replace('-', '_'), []) elif not cols and '_' in table: cols = declared_columns.get(table.replace('_', '-'), []) return cols def _hostvar_root(host_var: str, assignments: dict) -> str: """Trace a SQL WHERE host var through MOVE assignments to its source root. e.g. assignments['DBV-EMP-ID'] = [{'type': 'move', 'source_vars': ['R02EMP-ID']}] → root for DBV-EMP-ID is R02EMP-ID. Handles multi-hop MOVE chains. Returns the original host_var when no MOVE chain applies. """ seen = set() cur = host_var while cur in assignments and cur not in seen: seen.add(cur) al = assignments[cur] if isinstance(al, dict): al = [al] if not al: break last = al[-1] if last.get('type') == 'move' and last.get('source_vars'): cur = str(last['source_vars'][0]).upper() else: break return cur def _rec_get_ci(rec: dict, name: str): """Case-insensitive _rec_get. Returns the value string or None if not found.""" for scope in (rec, rec.get('working_storage', {}), rec.get('input', {})): for k, v in (scope or {}).items(): if str(k).upper() == name.upper(): return str(v) return None def _resolve_runtime_hostvar(rec: dict, host_var: str, assignments: dict) -> str: """Resolve a host var to the value the program actually uses at runtime. Returns only runtime-deterministic values: 1. The literal default in the root field's assignment chain when the root is fed by `ACCEPT ... FROM COMMAND-LINE` (e.g. `IF x = SPACES MOVE '202605' TO x`). The pipeline runs programs without command-line args, so that SPACES default is what the program actually uses at runtime (DBV-YEAR-MONTH -> ... -> WRK-YEAR-MONTH = '202605'). 2. The input-record field value at the end of a MOVE chain (e.g. `MOVE R02EMP-ID TO DBV-EMP-ID` -> R02EMP-ID's record value). Returns '' when no deterministic runtime value is derivable; the caller falls back to record heuristics. """ if not assignments: return '' root = _hostvar_root(host_var, assignments) al = assignments.get(root) if isinstance(al, dict): al = [al] if al and any(a.get('type') == 'accept' for a in al): for a in al: if a.get('type') == 'move_literal' and a.get('literal') is not None: return str(a['literal']) if root != host_var: rv = _rec_get_ci(rec, root) if rv is not None and str(rv).strip(): return str(rv) return '' def _resolve_where_hostvar(rec: dict, host_var: str, assignments: dict): """Resolve a WHERE host var to the value the program actually uses. If the host var is set by MOVE assignments (e.g. `MOVE R02EMP-ID TO DBV-EMP-ID`) before the SQL statement, trace to the input-record root and prefer its value, so the DB pre-seed matches the runtime query key. Falls back to the host var's own record value when no chain resolves. """ root = _hostvar_root(host_var, assignments) if root != host_var: rv = _rec_get_ci(rec, root) if rv is not None and str(rv).strip(): return str(rv) return _rec_get(rec, host_var, '') def _runtime_or_where_hostvar(rec: dict, host_var: str, assignments: dict) -> str: """Resolve a WHERE host var preferring the deterministic runtime value. `_resolve_runtime_hostvar` handles host vars fed by `ACCEPT ... FROM COMMAND-LINE` whose SPACES default the program substitutes at runtime (e.g. WRK-YEAR-MONTH = '202605'): the record carries a synthetic value ('I00001') that does NOT match the runtime query, so without this the seeded WHERE columns (EFFECTIVE-FROM/TO) fail the runtime predicate and the rows never load. Falls back to the existing MOVE-chain / record resolution when no deterministic runtime value exists. """ rv = _resolve_runtime_hostvar(rec, host_var, assignments) if rv: return rv return _resolve_where_hostvar(rec, host_var, assignments) def _resolve_between_operand(op: dict, rec: dict, assignments: dict): """Resolve a BETWEEN subject/lo/hi operand to a concrete string value.""" if not op: return None kind = op.get('kind') if kind == 'host_var': v = _resolve_where_hostvar(rec, op.get('name', ''), assignments) return v if kind == 'literal': return op.get('value') return None def _derive_runtime_hostvar(rec: dict, host_var: str, assignments: dict, fields: list) -> str | None: """Derive the runtime value of a WORKING-STORAGE host var via propagation. Records generated for coverage carry synthetic WRK-* values (e.g. WRK-TAXABLE-INCOME='000003701') that do NOT match what the program computes at runtime from the input keys (e.g. max(0, gross - deduction)). For a BETWEEN seed to match the runtime query key, we re-derive the value on a copy of the record: seed WORKING-STORAGE constants from their field VALUE clause, then run propagate_assignments. Returns None when the host var is not derivable (e.g. it is an input key with a direct MOVE chain). """ if not assignments or not fields: return None try: from .core import propagate_assignments probe = dict(rec) for f in fields: if f.get('value') is not None and f.get('section') == 'WORKING-STORAGE': probe.setdefault(f['name'], str(f['value'])) propagate_assignments(probe, assignments, fields) v = probe.get(host_var) if v is not None and str(v).strip(): return str(v) except Exception: pass return None def _between_seed_value(wc: dict, rec: dict, assignments: dict, col_name: str, fields: list = None) -> str | None: """Compute the DB column value for a BETWEEN where-constraint. Case A (`:hv BETWEEN LO-COL AND HI-COL`): the subject is a host var whose runtime value must fall inside the band; seed LO-COL / HI-COL with it. Case B (`SUBJ-COL BETWEEN :lo AND :hi`): the subject column is the one to seed; use the lo operand value so `lo <= col <= hi` holds. Returns None when this column is not governed by the BETWEEN constraint. """ if not wc or wc.get('type') != 'between' or wc.get('neg'): return None subject = wc.get('subject', {}) lo = wc.get('lo', {}) hi = wc.get('hi', {}) if subject.get('kind') == 'host_var': for side in (lo, hi): if side.get('kind') == 'column' and _norm_col(side.get('name', '')) == _norm_col(col_name): derived = _derive_runtime_hostvar(rec, subject.get('name', ''), assignments, fields or []) if derived: return derived v = _resolve_between_operand(subject, rec, assignments) if v and str(v).strip(): return str(v) elif subject.get('kind') == 'column' and _norm_col(subject.get('name', '')) == _norm_col(col_name): lo_val = _resolve_between_operand(lo, rec, assignments) hi_val = _resolve_between_operand(hi, rec, assignments) for v in (lo_val, hi_val): if v is not None and str(v).strip(): return str(v) return None def _input_pk_field(rec: dict, col_name: str) -> str | None: """Find the input-record field (R##-prefixed) that feeds a PK column. Only returns a field whose name ends with the hyphen-normalized column name. A YEAR-MONTH PK column therefore no longer falls back to the EMP-ID input field (which produced a truncated, wrong collision value). Returns None when no such input field exists (e.g. YEAR-MONTH comes from a WORKING-STORAGE ACCEPT default instead). """ if not rec: return None base = col_name.upper().replace('EMP_ID', 'EMP-ID') for key in rec: k = str(key).upper() if k.endswith(base) and k[:1] == 'R' and k[1:2].isdigit(): return key return None def _insert_collision_row(sql: dict, rec: dict, pk_cols: list[str], declared_columns: dict, assignments: dict = None) -> dict | None: """Build a DB pre-seed row whose PK collides with the input record's PK. For an INSERT statement, this creates a row with the same primary-key value as the value the program actually INSERTs at runtime, so the runtime INSERT fails with a duplicate-key SQL error (DB2 -803). Returns {} if no PK value could be derived from the record. """ if not rec or not pk_cols: return None columns = sql.get('columns', []) host_vars = sql.get('host_vars', []) if not columns or not host_vars: return None col_infos = declared_columns.get(sql.get('table', ''), []) if not col_infos: # Schema may declare the table with underscores while SQL uses hyphens. col_infos = declared_columns.get(sql.get('table', '').replace('-', '_'), []) if not col_infos: # Fall back to CHAR types inferred from the INSERT column list. col_infos = [{'name': c, 'db_type': 'CHAR', 'size': 20} for c in columns] col_to_hv = {} for c, hv in zip(columns, host_vars): col_to_hv[c.upper()] = hv col_to_hv[c.upper().replace('-', '_')] = hv row = {} for ci in col_infos: col_name = ci['name'].upper() col_name_alt = col_name.replace('-', '_') if '-' in col_name else col_name.replace('_', '-') if col_name not in pk_cols and col_name_alt not in pk_cols: continue pk_hit = col_name if col_name in pk_cols else col_name_alt hv = col_to_hv.get(col_name) or col_to_hv.get(col_name_alt) val = None # Prefer the runtime value the program actually INSERTs: # - EMP-ID <- MOVE chain to the input-record field (R02EMP-ID) # - YEAR-MONTH <- ACCEPT + IF-SPACES default literal ('202605') if hv: val = _resolve_runtime_hostvar(rec, hv, assignments or {}) if not val or not str(val).strip(): inp_field = _input_pk_field(rec, pk_hit) if inp_field: v = _rec_get(rec, inp_field, '') if v and str(v).strip(): val = str(v) if (not val or not str(val).strip()) and hv: val = _rec_get(rec, hv, '') if not val or not str(val).strip(): return None row[col_name] = _format_db_value(ci, str(val)) return row if row else None # ── 事务调度感知:按 88 级识别记录运行时执行的 SQL 种类 ── # 目的:UPDATE/DELETE 种子只对实际执行 UPDATE/DELETE 的记录建行, # 避免把 INSERT 记录的主键也预置进表(否则运行时 INSERT 全部 -803, # INSERT 成功分支不可达)。通用实现(88 级名 INSERT/UPDATE/DELETE 语义标记、 # 沿 MOVE 链追溯输入根字段),无程序名硬编码。 def _dispatch_field_and_sets(fields_dict): """返回 (调度字段名, {INSERT:[...], UPDATE:[...], DELETE:[...]}) 或 (None, {})。 调度字段 = 带 INSERT/UPDATE/DELETE 88 级语义标记的父字段(如 WRK-TRAN-TYPE)。 """ groups = {} for f in fields_dict or []: if isinstance(f, dict) and f.get('is_88') and f.get('parent'): groups.setdefault(f['parent'], []).append(f) for parent, kids in groups.items(): hi = any('INSERT' in str(k.get('name', '')).upper() for k in kids) hu = any('UPDATE' in str(k.get('name', '')).upper() for k in kids) hd = any('DELETE' in str(k.get('name', '')).upper() for k in kids) if hi and hu and hd: sets = {} for k in kids: nm = str(k.get('name', '')).upper() vals = k.get('values') or ([k.get('value')] if k.get('value') else []) group = next((g for g in ('INSERT', 'UPDATE', 'DELETE') if g in nm), None) if not group: continue for v in vals: sets.setdefault(group, set()).add(str(v).strip()) return parent, sets return None, {} def _trace_input_root(field_name, assignments): """沿 MOVE 链把调度字段追溯到输入记录根字段(WRK-TRAN-TYPE ← R01TRAN-TYPE)。""" seen = set() cur = field_name while cur and cur not in seen: seen.add(cur) al = assignments.get(cur) or [] if not al: break srcs = al[0].get('source_vars') or [] if len(srcs) != 1: break cur = srcs[0] return cur def _classify_record_dispatches(records, fields_dict, assignments): """返回 {path_idx: 'INSERT'|'UPDATE'|'DELETE'|'OTHER'}。 用记录的实际调度值(输入根字段,非合成工作区值)对照 88 级集合分类。 """ field, sets = _dispatch_field_and_sets(fields_dict) if not field: return {} root = _trace_input_root(field, assignments or {}) out = {} for i, r in enumerate(records or []): v = str(r.get(root, r.get(field, ''))).strip() out[i] = next((k for k in ('INSERT', 'UPDATE', 'DELETE') if v in sets.get(k, set())), 'OTHER') return out def build_db_input( branch_paths: list[tuple[list, dict]], fields_dict: list[dict], assignments: dict, sql_meta: list[dict], declared_columns: dict, records: list[dict] = None, insert_pk: dict[str, list[str]] = None, ) -> dict: """Generate DB input rows per branch path. Returns {table: [{col: val, ...}, ...]}. """ if not sql_meta: return {} db_input = {} seen_keys = {} seq_counter = itertools.count(1) # 事务调度感知:分类每条记录运行时执行的 SQL(INSERT/UPDATE/DELETE/OTHER)。 # UPDATE/DELETE 种子只对实际执行 UPDATE/DELETE 的记录建行,避免把 INSERT # 记录主键预置进表导致 INSERT 成功分支不可达。 dispatch = _classify_record_dispatches(records, fields_dict, assignments) # Downstream no-data coverage: when a path seeds rows for multiple SELECT # tables (e.g. EMP-MASTER then OVT-MONTHLY), drop the rows of the LAST # (runtime-order) table for the LAST SQL-ok path. That record then reaches # the downstream SELECT with no matching row (SQLCODE = 100), covering the # downstream SELECT's "no data" branch (IF SQLCODE = 0 ELSE path). sql_select_tables = [m.get('table') for m in sql_meta if m.get('type') == 'exec_sql_select' and m.get('table')] unique_select_tables = [] for t in sql_select_tables: if t not in unique_select_tables: unique_select_tables.append(t) drop_table = unique_select_tables[-1] if len(unique_select_tables) >= 2 else None sql_ok_paths = [i for i, (pc, _pa) in enumerate(branch_paths) if _path_has_sql_ok(pc)] last_sql_ok_path = sql_ok_paths[-1] if len(sql_ok_paths) >= 2 else -1 # Collect all SQL meta per path for path_idx, (path_cons, path_assign) in enumerate(branch_paths): sql_ok = _path_has_sql_ok(path_cons) rec = records[path_idx] if records and path_idx < len(records) else {} if path_idx == 0 and rec: with open(r'C:\Users\marye\AppData\Local\Temp\opencode\build_db_input_debug.txt', 'w') as _f: _f.write(f"rec keys count={len(rec)}\n") _f.write(f"has_HV-ANNUAL-H={'HV-ANNUAL-H' in rec}\n") _f.write(f"has_working_storage={'working_storage' in rec}\n") _f.write(f"HV-ANNUAL-H via _rec_get={_rec_get(rec, 'HV-ANNUAL-H', 'NOT_FOUND')!r}\n") if 'HV-ANNUAL-H' in rec: _f.write(f"HV-ANNUAL-H value={rec['HV-ANNUAL-H']!r}\n") _f.write(f"all keys sorted={sorted(rec.keys())}\n") for sql in sql_meta: atype = sql.get('type', '') table = sql.get('table', '') if not table: continue where_cons = sql.get('where_constraints', []) if table not in db_input: db_input[table] = [] seen_keys[table] = set() if atype == 'exec_sql_insert': # INSERT creates rows at runtime; no initial rows needed, # EXCEPT a PK-collision row so the duplicate-key error path # (SQLCODE <> 0 / -803) is reachable at runtime. if insert_pk: pk_cols = [c.upper() for c in insert_pk.get(table, [])] if pk_cols and not _path_has_sql_ok(path_cons): # 碰撞行必须命中真实执行 INSERT 的记录:若本路径记录并非 # INSERT 类(如合成记录走了 WHEN OTHER),回退到首条 # INSERT 类记录的主键,确保运行时 -803 分支可达。 target = rec if dispatch.get(path_idx) == 'INSERT' else None if target is None: ins_idx = next((i for i in sorted(dispatch) if dispatch[i] == 'INSERT'), None) target = records[ins_idx] if ins_idx is not None else rec row = _insert_collision_row(sql, target, pk_cols, declared_columns, assignments=assignments) if row: db_input[table].append(row) continue if atype in ('exec_sql_delete', 'exec_sql_update'): # DELETE/UPDATE needs existing rows to act on if not sql_ok: continue # 仅对实际执行 UPDATE/DELETE 的记录建行(避免污染 INSERT/OTHER 主键)。 # 无 88 级调度字段可识别时不做门控(回退旧行为)。 if dispatch and dispatch.get(path_idx) not in ('UPDATE', 'DELETE'): continue col_infos = _declared_cols_for(declared_columns, table) if not col_infos: col_infos = _infer_columns_from_where(where_cons) row = {} for ci in col_infos: col_name = ci['name'].upper() val = None for wc in where_cons: wc_col = wc.get('col', '').upper().split('.')[-1] if _norm_col(wc_col) != _norm_col(col_name): continue if wc['type'] == 'literal': val = wc.get('literal', '') break elif wc['type'] == 'host_var': hv = wc.get('host_var', '').upper() val = _runtime_or_where_hostvar(rec, hv, assignments) break if val is None or not val.strip(): val = _rec_get(rec, ci['name'], '') if val and val.strip(): row[ci['name']] = _format_db_value(ci, val) if not row: row['_path'] = str(path_idx) db_input[table].append(row) continue # exec_sql_select (and any future read-only types) if not sql_ok: continue row = {} col_infos = _declared_cols_for(declared_columns, table) if not col_infos: col_infos = _infer_columns_from_where(where_cons) into_vars = sql.get('into_vars', []) # Map INTO vars → actual SQL column names from SELECT clause select_cols = _parse_select_columns(sql.get('select_list', '')) into_to_col: dict[str, str] = {} for i, iv in enumerate(into_vars): if i < len(select_cols): into_to_col[iv] = select_cols[i] # Use SQL column names (not INTO var names) for row keys for sc in select_cols: if not any(_norm_col(sc) == _norm_col(c['name']) for c in col_infos): col_infos.append({'name': sc, 'db_type': 'CHAR', 'size': 20}) where_cols = set() for col_info in col_infos: col_name = col_info['name'] val = None for wc in where_cons: if wc['type'] == 'between': bv = _between_seed_value(wc, rec, assignments, col_name, fields_dict) if bv is not None: val = bv where_cols.add(col_name) break if wc['type'] == 'literal' and _norm_col(wc.get('col', '')) == _norm_col(col_name): val = wc.get('literal', '') where_cols.add(col_name) break if wc['type'] == 'host_var': wc_col = wc.get('col', '').upper().split('.')[-1] if _norm_col(wc_col) != _norm_col(col_name): continue hv = wc.get('host_var', '').upper() for pc_field, pc_op, pc_val, pc_want in path_cons: if pc_field == hv: val = pc_val if pc_want else '' where_cols.add(col_name) break if val is None and _rec_has(rec, hv): val = _runtime_or_where_hostvar(rec, hv, assignments) where_cols.add(col_name) # Try to find value from INTO variable in the record if val is None: for iv, scola in into_to_col.items(): if _norm_col(scola) == _norm_col(col_name) and _rec_has(rec, iv): val = _rec_get(rec, iv, '') break # Try COBOL field name mapping if val is None: cobol_field = guess_cobol_field(col_name, table, declared_columns) if _rec_has(rec, cobol_field): val = _rec_get(rec, cobol_field, '') if val is not None: row[col_name] = _format_db_value(col_info, val) else: row[col_name] = _format_db_value(col_info, str(next(seq_counter))) if not row: row['_path'] = str(path_idx) if col_infos: first_col = col_infos[0]['name'] if first_col in row and (not where_cols or first_col not in where_cols): row[first_col] = _make_key_unique(row[first_col], path_idx, seen_keys[table]) if drop_table is not None and table == drop_table and path_idx == last_sql_ok_path: continue db_input[table].append(row) return db_input