1#!/usr/bin/env python3
2# -*- coding: utf-8 -*-
3
4# Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
5# See https://llvm.org/LICENSE.txt for license information.
6# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
7
8# Script for updating SPIR-V dialect by scraping information from SPIR-V
9# HTML and JSON specs from the Internet.
10#
11# For example, to define the enum attribute for SPIR-V memory model:
12#
13# ./gen_spirv_dialect.py --base_td_path /path/to/SPIRVBase.td \
14#                        --new-enum MemoryModel
15#
16# The 'operand_kinds' dict of spirv.core.grammar.json contains all supported
17# SPIR-V enum classes.
18
19import itertools
20import re
21import requests
22import textwrap
23
24SPIRV_HTML_SPEC_URL = 'https://www.khronos.org/registry/spir-v/specs/unified1/SPIRV.html'
25SPIRV_JSON_SPEC_URL = 'https://raw.githubusercontent.com/KhronosGroup/SPIRV-Headers/master/include/spirv/unified1/spirv.core.grammar.json'
26
27AUTOGEN_OP_DEF_SEPARATOR = '\n// -----\n\n'
28AUTOGEN_ENUM_SECTION_MARKER = 'enum section. Generated from SPIR-V spec; DO NOT MODIFY!'
29AUTOGEN_OPCODE_SECTION_MARKER = (
30    'opcode section. Generated from SPIR-V spec; DO NOT MODIFY!')
31
32
33def get_spirv_doc_from_html_spec():
34  """Extracts instruction documentation from SPIR-V HTML spec.
35
36  Returns:
37    - A dict mapping from instruction opcode to documentation.
38  """
39  response = requests.get(SPIRV_HTML_SPEC_URL)
40  spec = response.content
41
42  from bs4 import BeautifulSoup
43  spirv = BeautifulSoup(spec, 'html.parser')
44
45  section_anchor = spirv.find('h3', {'id': '_a_id_instructions_a_instructions'})
46
47  doc = {}
48
49  for section in section_anchor.parent.find_all('div', {'class': 'sect3'}):
50    for table in section.find_all('table'):
51      inst_html = table.tbody.tr.td.p
52      opname = inst_html.a['id']
53      # Ignore the first line, which is just the opname.
54      doc[opname] = inst_html.text.split('\n', 1)[1].strip()
55
56  return doc
57
58
59def get_spirv_grammar_from_json_spec():
60  """Extracts operand kind and instruction grammar from SPIR-V JSON spec.
61
62  Returns:
63    - A list containing all operand kinds' grammar
64    - A list containing all instructions' grammar
65  """
66  response = requests.get(SPIRV_JSON_SPEC_URL)
67  spec = response.content
68
69  import json
70  spirv = json.loads(spec)
71
72  return spirv['operand_kinds'], spirv['instructions']
73
74
75def split_list_into_sublists(items, offset):
76  """Split the list of items into multiple sublists.
77
78  This is to make sure the string composed from each sublist won't exceed
79  80 characters.
80
81  Arguments:
82    - items: a list of strings
83    - offset: the offset in calculating each sublist's length
84  """
85  chuncks = []
86  chunk = []
87  chunk_len = 0
88
89  for item in items:
90    chunk_len += len(item) + 2
91    if chunk_len > 80:
92      chuncks.append(chunk)
93      chunk = []
94      chunk_len = len(item) + 2
95    chunk.append(item)
96
97  if len(chunk) != 0:
98    chuncks.append(chunk)
99
100  return chuncks
101
102
103def uniquify_enum_cases(lst):
104  """Prunes duplicate enum cases from the list.
105
106  Arguments:
107   - lst: List whose elements are to be uniqued. Assumes each element is a
108     (symbol, value) pair and elements already sorted according to value.
109
110  Returns:
111   - A list with all duplicates removed. The elements are sorted according to
112     value and, for each value, uniqued according to symbol.
113     original list,
114   - A map from deduplicated cases to the uniqued case.
115  """
116  cases = lst
117  uniqued_cases = []
118  duplicated_cases = {}
119
120  # First sort according to the value
121  cases.sort(key=lambda x: x[1])
122
123  # Then group them according to the value
124  for _, groups in itertools.groupby(cases, key=lambda x: x[1]):
125    # For each value, sort according to the enumerant symbol.
126    sorted_group = sorted(groups, key=lambda x: x[0])
127    # Keep the "smallest" case, which is typically the symbol without extension
128    # suffix. But we have special cases that we want to fix.
129    case = sorted_group[0]
130    for i in range(1, len(sorted_group)):
131      duplicated_cases[sorted_group[i][0]] = case[0]
132    if case[0] == 'HlslSemanticGOOGLE':
133      assert len(sorted_group) == 2, 'unexpected new variant for HlslSemantic'
134      case = sorted_group[1]
135      duplicated_cases[sorted_group[0][0]] = case[0]
136    uniqued_cases.append(case)
137
138  return uniqued_cases, duplicated_cases
139
140
141def toposort(dag, sort_fn):
142  """Topologically sorts the given dag.
143
144  Arguments:
145    - dag: a dict mapping from a node to its incoming nodes.
146    - sort_fn: a function for sorting nodes in the same batch.
147
148  Returns:
149    A list containing topologically sorted nodes.
150  """
151
152  # Returns the next batch of nodes without incoming edges
153  def get_next_batch(dag):
154    while True:
155      no_prev_nodes = set(node for node, prev in dag.items() if not prev)
156      if not no_prev_nodes:
157        break
158      yield sorted(no_prev_nodes, key=sort_fn)
159      dag = {
160          node: (prev - no_prev_nodes)
161          for node, prev in dag.items()
162          if node not in no_prev_nodes
163      }
164    assert not dag, 'found cyclic dependency'
165
166  sorted_nodes = []
167  for batch in get_next_batch(dag):
168    sorted_nodes.extend(batch)
169
170  return sorted_nodes
171
172
173def toposort_capabilities(all_cases, capability_mapping):
174  """Returns topologically sorted capability (symbol, value) pairs.
175
176  Arguments:
177    - all_cases: all capability cases (containing symbol, value, and implied
178      capabilities).
179    - capability_mapping: mapping from duplicated capability symbols to the
180      canonicalized symbol chosen for SPIRVBase.td.
181
182  Returns:
183    A list containing topologically sorted capability (symbol, value) pairs.
184  """
185  dag = {}
186  name_to_value = {}
187  for case in all_cases:
188    # Get the current capability.
189    cur = case['enumerant']
190    name_to_value[cur] = case['value']
191    # Ignore duplicated symbols.
192    if cur in capability_mapping:
193      continue
194
195    # Get capabilities implied by the current capability.
196    prev = case.get('capabilities', [])
197    uniqued_prev = set([capability_mapping.get(c, c) for c in prev])
198    dag[cur] = uniqued_prev
199
200  sorted_caps = toposort(dag, lambda x: name_to_value[x])
201  # Attach the capability's value as the second component of the pair.
202  return [(c, name_to_value[c]) for c in sorted_caps]
203
204
205def get_capability_mapping(operand_kinds):
206  """Returns the capability mapping from duplicated cases to canonicalized ones.
207
208  Arguments:
209    - operand_kinds: all operand kinds' grammar spec
210
211  Returns:
212    - A map mapping from duplicated capability symbols to the canonicalized
213      symbol chosen for SPIRVBase.td.
214  """
215  # Find the operand kind for capability
216  cap_kind = {}
217  for kind in operand_kinds:
218    if kind['kind'] == 'Capability':
219      cap_kind = kind
220
221  kind_cases = [
222      (case['enumerant'], case['value']) for case in cap_kind['enumerants']
223  ]
224  _, capability_mapping = uniquify_enum_cases(kind_cases)
225
226  return capability_mapping
227
228
229def get_availability_spec(enum_case, capability_mapping, for_op, for_cap):
230  """Returns the availability specification string for the given enum case.
231
232  Arguments:
233    - enum_case: the enum case to generate availability spec for. It may contain
234      'version', 'lastVersion', 'extensions', or 'capabilities'.
235    - capability_mapping: mapping from duplicated capability symbols to the
236      canonicalized symbol chosen for SPIRVBase.td.
237    - for_op: bool value indicating whether this is the availability spec for an
238      op itself.
239    - for_cap: bool value indicating whether this is the availability spec for
240      capabilities themselves.
241
242  Returns:
243    - A `let availability = [...];` string if with availability spec or
244      empty string if without availability spec
245  """
246  assert not (for_op and for_cap), 'cannot set both for_op and for_cap'
247
248  DEFAULT_MIN_VERSION = 'MinVersion<SPV_V_1_0>'
249  DEFAULT_MAX_VERSION = 'MaxVersion<SPV_V_1_5>'
250  DEFAULT_CAP = 'Capability<[]>'
251  DEFAULT_EXT = 'Extension<[]>'
252
253  min_version = enum_case.get('version', '')
254  if min_version == 'None':
255    min_version = ''
256  elif min_version:
257    min_version = 'MinVersion<SPV_V_{}>'.format(min_version.replace('.', '_'))
258  # TODO(antiagainst): delete this once ODS can support dialect-specific content
259  # and we can use omission to mean no requirements.
260  if for_op and not min_version:
261    min_version = DEFAULT_MIN_VERSION
262
263  max_version = enum_case.get('lastVersion', '')
264  if max_version:
265    max_version = 'MaxVersion<SPV_V_{}>'.format(max_version.replace('.', '_'))
266  # TODO(antiagainst): delete this once ODS can support dialect-specific content
267  # and we can use omission to mean no requirements.
268  if for_op and not max_version:
269    max_version = DEFAULT_MAX_VERSION
270
271  exts = enum_case.get('extensions', [])
272  if exts:
273    exts = 'Extension<[{}]>'.format(', '.join(sorted(set(exts))))
274    # We need to strip the minimal version requirement if this symbol is
275    # available via an extension, which means *any* SPIR-V version can support
276    # it as long as the extension is provided. The grammar's 'version' field
277    # under such case should be interpreted as this symbol is introduced as
278    # a core symbol since the given version, rather than a minimal version
279    # requirement.
280    min_version = DEFAULT_MIN_VERSION if for_op else ''
281  # TODO(antiagainst): delete this once ODS can support dialect-specific content
282  # and we can use omission to mean no requirements.
283  if for_op and not exts:
284    exts = DEFAULT_EXT
285
286  caps = enum_case.get('capabilities', [])
287  implies = ''
288  if caps:
289    canonicalized_caps = []
290    for c in caps:
291      if c in capability_mapping:
292        canonicalized_caps.append(capability_mapping[c])
293      else:
294        canonicalized_caps.append(c)
295    prefixed_caps = [
296        'SPV_C_{}'.format(c) for c in sorted(set(canonicalized_caps))
297    ]
298    if for_cap:
299      # If this is generating the availability for capabilities, we need to
300      # put the capability "requirements" in implies field because now
301      # the "capabilities" field in the source grammar means so.
302      caps = ''
303      implies = 'list<I32EnumAttrCase> implies = [{}];'.format(
304          ', '.join(prefixed_caps))
305    else:
306      caps = 'Capability<[{}]>'.format(', '.join(prefixed_caps))
307      implies = ''
308  # TODO(antiagainst): delete this once ODS can support dialect-specific content
309  # and we can use omission to mean no requirements.
310  if for_op and not caps:
311    caps = DEFAULT_CAP
312
313  avail = ''
314  # Compose availability spec if any of the requirements is not empty.
315  # For ops, because we have a default in SPV_Op class, omit if the spec
316  # is the same.
317  if (min_version or max_version or caps or exts) and not (
318      for_op and min_version == DEFAULT_MIN_VERSION and
319      max_version == DEFAULT_MAX_VERSION and caps == DEFAULT_CAP and
320      exts == DEFAULT_EXT):
321    joined_spec = ',\n    '.join(
322        [e for e in [min_version, max_version, exts, caps] if e])
323    avail = '{} availability = [\n    {}\n  ];'.format(
324        'let' if for_op else 'list<Availability>', joined_spec)
325
326  return '{}{}{}'.format(implies, '\n  ' if implies and avail else '', avail)
327
328
329def gen_operand_kind_enum_attr(operand_kind, capability_mapping):
330  """Generates the TableGen EnumAttr definition for the given operand kind.
331
332  Returns:
333    - The operand kind's name
334    - A string containing the TableGen EnumAttr definition
335  """
336  if 'enumerants' not in operand_kind:
337    return '', ''
338
339  # Returns a symbol for the given case in the given kind. This function
340  # handles Dim specially to avoid having numbers as the start of symbols,
341  # which does not play well with C++ and the MLIR parser.
342  def get_case_symbol(kind_name, case_name):
343    if kind_name == 'Dim':
344      if case_name == '1D' or case_name == '2D' or case_name == '3D':
345        return 'Dim{}'.format(case_name)
346    return case_name
347
348  kind_name = operand_kind['kind']
349  is_bit_enum = operand_kind['category'] == 'BitEnum'
350  kind_category = 'Bit' if is_bit_enum else 'I32'
351  kind_acronym = ''.join([c for c in kind_name if c >= 'A' and c <= 'Z'])
352
353  name_to_case_dict = {}
354  for case in operand_kind['enumerants']:
355    name_to_case_dict[case['enumerant']] = case
356
357  if kind_name == 'Capability':
358    # Special treatment for capability cases: we need to sort them topologically
359    # because a capability can refer to another via the 'implies' field.
360    kind_cases = toposort_capabilities(operand_kind['enumerants'],
361                                       capability_mapping)
362  else:
363    kind_cases = [(case['enumerant'], case['value'])
364                  for case in operand_kind['enumerants']]
365    kind_cases, _ = uniquify_enum_cases(kind_cases)
366  max_len = max([len(symbol) for (symbol, _) in kind_cases])
367
368  # Generate the definition for each enum case
369  fmt_str = 'def SPV_{acronym}_{case} {colon:>{offset}} '\
370            '{category}EnumAttrCase<"{symbol}", {value}>{avail}'
371  case_defs = []
372  for case in kind_cases:
373    avail = get_availability_spec(name_to_case_dict[case[0]],
374                                  capability_mapping,
375                                  False, kind_name == 'Capability')
376    case_def = fmt_str.format(
377        category=kind_category,
378        acronym=kind_acronym,
379        case=case[0],
380        symbol=get_case_symbol(kind_name, case[0]),
381        value=case[1],
382        avail=' {{\n  {}\n}}'.format(avail) if avail else ';',
383        colon=':',
384        offset=(max_len + 1 - len(case[0])))
385    case_defs.append(case_def)
386  case_defs = '\n'.join(case_defs)
387
388  # Generate the list of enum case names
389  fmt_str = 'SPV_{acronym}_{symbol}';
390  case_names = [fmt_str.format(acronym=kind_acronym,symbol=case[0])
391                for case in kind_cases]
392
393  # Split them into sublists and concatenate into multiple lines
394  case_names = split_list_into_sublists(case_names, 6)
395  case_names = ['{:6}'.format('') + ', '.join(sublist)
396                for sublist in case_names]
397  case_names = ',\n'.join(case_names)
398
399  # Generate the enum attribute definition
400  enum_attr = '''def SPV_{name}Attr :
401    SPV_{category}EnumAttr<"{name}", "valid SPIR-V {name}", [
402{cases}
403    ]>;'''.format(
404          name=kind_name, category=kind_category, cases=case_names)
405  return kind_name, case_defs + '\n\n' + enum_attr
406
407
408def gen_opcode(instructions):
409  """ Generates the TableGen definition to map opname to opcode
410
411  Returns:
412    - A string containing the TableGen SPV_OpCode definition
413  """
414
415  max_len = max([len(inst['opname']) for inst in instructions])
416  def_fmt_str = 'def SPV_OC_{name} {colon:>{offset}} '\
417            'I32EnumAttrCase<"{name}", {value}>;'
418  opcode_defs = [
419      def_fmt_str.format(
420          name=inst['opname'],
421          value=inst['opcode'],
422          colon=':',
423          offset=(max_len + 1 - len(inst['opname']))) for inst in instructions
424  ]
425  opcode_str = '\n'.join(opcode_defs)
426
427  decl_fmt_str = 'SPV_OC_{name}'
428  opcode_list = [
429      decl_fmt_str.format(name=inst['opname']) for inst in instructions
430  ]
431  opcode_list = split_list_into_sublists(opcode_list, 6)
432  opcode_list = [
433      '{:6}'.format('') + ', '.join(sublist) for sublist in opcode_list
434  ]
435  opcode_list = ',\n'.join(opcode_list)
436  enum_attr = 'def SPV_OpcodeAttr :\n'\
437              '    SPV_I32EnumAttr<"{name}", "valid SPIR-V instructions", [\n'\
438              '{lst}\n'\
439              '    ]>;'.format(name='Opcode', lst=opcode_list)
440  return opcode_str + '\n\n' + enum_attr
441
442
443def update_td_opcodes(path, instructions, filter_list):
444  """Updates SPIRBase.td with new generated opcode cases.
445
446  Arguments:
447    - path: the path to SPIRBase.td
448    - instructions: a list containing all SPIR-V instructions' grammar
449    - filter_list: a list containing new opnames to add
450  """
451
452  with open(path, 'r') as f:
453    content = f.read()
454
455  content = content.split(AUTOGEN_OPCODE_SECTION_MARKER)
456  assert len(content) == 3
457
458  # Extend opcode list with existing list
459  existing_opcodes = [k[11:] for k in re.findall('def SPV_OC_\w+', content[1])]
460  filter_list.extend(existing_opcodes)
461  filter_list = list(set(filter_list))
462
463  # Generate the opcode for all instructions in SPIR-V
464  filter_instrs = list(
465      filter(lambda inst: (inst['opname'] in filter_list), instructions))
466  # Sort instruction based on opcode
467  filter_instrs.sort(key=lambda inst: inst['opcode'])
468  opcode = gen_opcode(filter_instrs)
469
470  # Substitute the opcode
471  content = content[0] + AUTOGEN_OPCODE_SECTION_MARKER + '\n\n' + \
472        opcode + '\n\n// End ' + AUTOGEN_OPCODE_SECTION_MARKER \
473        + content[2]
474
475  with open(path, 'w') as f:
476    f.write(content)
477
478
479def update_td_enum_attrs(path, operand_kinds, filter_list):
480  """Updates SPIRBase.td with new generated enum definitions.
481
482  Arguments:
483    - path: the path to SPIRBase.td
484    - operand_kinds: a list containing all operand kinds' grammar
485    - filter_list: a list containing new enums to add
486  """
487  with open(path, 'r') as f:
488    content = f.read()
489
490  content = content.split(AUTOGEN_ENUM_SECTION_MARKER)
491  assert len(content) == 3
492
493  # Extend filter list with existing enum definitions
494  existing_kinds = [
495      k[8:-4] for k in re.findall('def SPV_\w+Attr', content[1])]
496  filter_list.extend(existing_kinds)
497
498  capability_mapping = get_capability_mapping(operand_kinds)
499
500  # Generate definitions for all enums in filter list
501  defs = [
502      gen_operand_kind_enum_attr(kind, capability_mapping)
503      for kind in operand_kinds
504      if kind['kind'] in filter_list
505  ]
506  # Sort alphabetically according to enum name
507  defs.sort(key=lambda enum : enum[0])
508  # Only keep the definitions from now on
509  # Put Capability's definition at the very beginning because capability cases
510  # will be referenced later
511  defs = [enum[1] for enum in defs if enum[0] == 'Capability'
512         ] + [enum[1] for enum in defs if enum[0] != 'Capability']
513
514  # Substitute the old section
515  content = content[0] + AUTOGEN_ENUM_SECTION_MARKER + '\n\n' + \
516      '\n\n'.join(defs) + "\n\n// End " + AUTOGEN_ENUM_SECTION_MARKER  \
517      + content[2];
518
519  with open(path, 'w') as f:
520    f.write(content)
521
522
523def snake_casify(name):
524  """Turns the given name to follow snake_case convention."""
525  name = re.sub('\W+', '', name).split()
526  name = [s.lower() for s in name]
527  return '_'.join(name)
528
529
530def map_spec_operand_to_ods_argument(operand):
531  """Maps an operand in SPIR-V JSON spec to an op argument in ODS.
532
533  Arguments:
534    - A dict containing the operand's kind, quantifier, and name
535
536  Returns:
537    - A string containing both the type and name for the argument
538  """
539  kind = operand['kind']
540  quantifier = operand.get('quantifier', '')
541
542  # These instruction "operands" are for encoding the results; they should
543  # not be handled here.
544  assert kind != 'IdResultType', 'unexpected to handle "IdResultType" kind'
545  assert kind != 'IdResult', 'unexpected to handle "IdResult" kind'
546
547  if kind == 'IdRef':
548    if quantifier == '':
549      arg_type = 'SPV_Type'
550    elif quantifier == '?':
551      arg_type = 'SPV_Optional<SPV_Type>'
552    else:
553      arg_type = 'Variadic<SPV_Type>'
554  elif kind == 'IdMemorySemantics' or kind == 'IdScope':
555    # TODO(antiagainst): Need to further constrain 'IdMemorySemantics'
556    # and 'IdScope' given that they should be generated from OpConstant.
557    assert quantifier == '', ('unexpected to have optional/variadic memory '
558                              'semantics or scope <id>')
559    arg_type = 'SPV_' + kind[2:] + 'Attr'
560  elif kind == 'LiteralInteger':
561    if quantifier == '':
562      arg_type = 'I32Attr'
563    elif quantifier == '?':
564      arg_type = 'OptionalAttr<I32Attr>'
565    else:
566      arg_type = 'OptionalAttr<I32ArrayAttr>'
567  elif kind == 'LiteralString' or \
568      kind == 'LiteralContextDependentNumber' or \
569      kind == 'LiteralExtInstInteger' or \
570      kind == 'LiteralSpecConstantOpInteger' or \
571      kind == 'PairLiteralIntegerIdRef' or \
572      kind == 'PairIdRefLiteralInteger' or \
573      kind == 'PairIdRefIdRef':
574    assert False, '"{}" kind unimplemented'.format(kind)
575  else:
576    # The rest are all enum operands that we represent with op attributes.
577    assert quantifier != '*', 'unexpected to have variadic enum attribute'
578    arg_type = 'SPV_{}Attr'.format(kind)
579    if quantifier == '?':
580      arg_type = 'OptionalAttr<{}>'.format(arg_type)
581
582  name = operand.get('name', '')
583  name = snake_casify(name) if name else kind.lower()
584
585  return '{}:${}'.format(arg_type, name)
586
587
588def get_description(text, appendix):
589  """Generates the description for the given SPIR-V instruction.
590
591  Arguments:
592    - text: Textual description of the operation as string.
593    - appendix: Additional contents to attach in description as string,
594                includking IR examples, and others.
595
596  Returns:
597    - A string that corresponds to the description of the Tablegen op.
598  """
599  fmt_str = '{text}\n\n    <!-- End of AutoGen section -->\n{appendix}\n  '
600  return fmt_str.format(text=text, appendix=appendix)
601
602
603def get_op_definition(instruction, doc, existing_info, capability_mapping):
604  """Generates the TableGen op definition for the given SPIR-V instruction.
605
606  Arguments:
607    - instruction: the instruction's SPIR-V JSON grammar
608    - doc: the instruction's SPIR-V HTML doc
609    - existing_info: a dict containing potential manually specified sections for
610      this instruction
611    - capability_mapping: mapping from duplicated capability symbols to the
612                   canonicalized symbol chosen for SPIRVBase.td
613
614  Returns:
615    - A string containing the TableGen op definition
616  """
617  fmt_str = ('def SPV_{opname}Op : '
618             'SPV_{inst_category}<"{opname}"{category_args}[{traits}]> '
619             '{{\n  let summary = {summary};\n\n  let description = '
620             '[{{\n{description}}}];{availability}\n')
621  inst_category = existing_info.get('inst_category', 'Op')
622  if inst_category == 'Op':
623    fmt_str +='\n  let arguments = (ins{args});\n\n'\
624              '  let results = (outs{results});\n'
625
626  fmt_str +='{extras}'\
627            '}}\n'
628
629  opname = instruction['opname'][2:]
630  category_args = existing_info.get('category_args', '')
631
632  if '\n' in doc:
633    summary, text = doc.split('\n', 1)
634  else:
635    summary = doc
636    text = ''
637  wrapper = textwrap.TextWrapper(
638      width=76, initial_indent='    ', subsequent_indent='    ')
639
640  # Format summary. If the summary can fit in the same line, we print it out
641  # as a "-quoted string; otherwise, wrap the lines using "[{...}]".
642  summary = summary.strip();
643  if len(summary) + len('  let summary = "";') <= 80:
644    summary = '"{}"'.format(summary)
645  else:
646    summary = '[{{\n{}\n  }}]'.format(wrapper.fill(summary))
647
648  # Wrap text
649  text = text.split('\n')
650  text = [wrapper.fill(line) for line in text if line]
651  text = '\n\n'.join(text)
652
653  operands = instruction.get('operands', [])
654
655  # Op availability
656  avail = ''
657  # We assume other instruction categories has a base availability spec, so
658  # only add this if this is directly using SPV_Op as the base.
659  if inst_category == 'Op':
660    avail = get_availability_spec(instruction, capability_mapping, True, False)
661    if avail:
662      avail = '\n\n  {0}'.format(avail)
663
664  # Set op's result
665  results = ''
666  if len(operands) > 0 and operands[0]['kind'] == 'IdResultType':
667    results = '\n    SPV_Type:$result\n  '
668    operands = operands[1:]
669  if 'results' in existing_info:
670    results = existing_info['results']
671
672  # Ignore the operand standing for the result <id>
673  if len(operands) > 0 and operands[0]['kind'] == 'IdResult':
674    operands = operands[1:]
675
676  # Set op' argument
677  arguments = existing_info.get('arguments', None)
678  if arguments is None:
679    arguments = [map_spec_operand_to_ods_argument(o) for o in operands]
680    arguments = ',\n    '.join(arguments)
681    if arguments:
682      # Prepend and append whitespace for formatting
683      arguments = '\n    {}\n  '.format(arguments)
684
685  description = existing_info.get('description', None)
686  if description is None:
687    assembly = '\n    ```\n'\
688               '    [TODO]\n'\
689               '    ```mlir\n\n'\
690               '    #### Example:\n\n'\
691               '    ```\n'\
692               '    [TODO]\n' \
693               '    ```'
694    description = get_description(text, assembly)
695
696  return fmt_str.format(
697      opname=opname,
698      category_args=category_args,
699      inst_category=inst_category,
700      traits=existing_info.get('traits', ''),
701      summary=summary,
702      description=description,
703      availability=avail,
704      args=arguments,
705      results=results,
706      extras=existing_info.get('extras', ''))
707
708
709def get_string_between(base, start, end):
710  """Extracts a substring with a specified start and end from a string.
711
712  Arguments:
713    - base: string to extract from.
714    - start: string to use as the start of the substring.
715    - end: string to use as the end of the substring.
716
717  Returns:
718    - The substring if found
719    - The part of the base after end of the substring. Is the base string itself
720      if the substring wasnt found.
721  """
722  split = base.split(start, 1)
723  if len(split) == 2:
724    rest = split[1].split(end, 1)
725    assert len(rest) == 2, \
726           'cannot find end "{end}" while extracting substring '\
727           'starting with {start}'.format(start=start, end=end)
728    return rest[0].rstrip(end), rest[1]
729  return '', split[0]
730
731
732def get_string_between_nested(base, start, end):
733  """Extracts a substring with a nested start and end from a string.
734
735  Arguments:
736    - base: string to extract from.
737    - start: string to use as the start of the substring.
738    - end: string to use as the end of the substring.
739
740  Returns:
741    - The substring if found
742    - The part of the base after end of the substring. Is the base string itself
743      if the substring wasn't found.
744  """
745  split = base.split(start, 1)
746  if len(split) == 2:
747    # Handle nesting delimiters
748    rest = split[1]
749    unmatched_start = 1
750    index = 0
751    while unmatched_start > 0 and index < len(rest):
752      if rest[index:].startswith(end):
753        unmatched_start -= 1
754        if unmatched_start == 0:
755          break
756        index += len(end)
757      elif rest[index:].startswith(start):
758        unmatched_start += 1
759        index += len(start)
760      else:
761        index += 1
762
763    assert index < len(rest), \
764           'cannot find end "{end}" while extracting substring '\
765           'starting with "{start}"'.format(start=start, end=end)
766    return rest[:index], rest[index + len(end):]
767  return '', split[0]
768
769
770def extract_td_op_info(op_def):
771  """Extracts potentially manually specified sections in op's definition.
772
773  Arguments: - A string containing the op's TableGen definition
774    - doc: the instruction's SPIR-V HTML doc
775
776  Returns:
777    - A dict containing potential manually specified sections
778  """
779  # Get opname
780  opname = [o[8:-2] for o in re.findall('def SPV_\w+Op', op_def)]
781  assert len(opname) == 1, 'more than one ops in the same section!'
782  opname = opname[0]
783
784  # Get instruction category
785  inst_category = [
786      o[4:] for o in re.findall('SPV_\w+Op',
787                                op_def.split(':', 1)[1])
788  ]
789  assert len(inst_category) <= 1, 'more than one ops in the same section!'
790  inst_category = inst_category[0] if len(inst_category) == 1 else 'Op'
791
792  # Get category_args
793  op_tmpl_params, _ = get_string_between_nested(op_def, '<', '>')
794  opstringname, rest = get_string_between(op_tmpl_params, '"', '"')
795  category_args = rest.split('[', 1)[0]
796
797  # Get traits
798  traits, _ = get_string_between_nested(rest, '[', ']')
799
800  # Get description
801  description, rest = get_string_between(op_def, 'let description = [{\n',
802                                         '}];\n')
803
804  # Get arguments
805  args, rest = get_string_between(rest, '  let arguments = (ins', ');\n')
806
807  # Get results
808  results, rest = get_string_between(rest, '  let results = (outs', ');\n')
809
810  extras = rest.strip(' }\n')
811  if extras:
812    extras = '\n  {}\n'.format(extras)
813
814  return {
815      # Prefix with 'Op' to make it consistent with SPIR-V spec
816      'opname': 'Op{}'.format(opname),
817      'inst_category': inst_category,
818      'category_args': category_args,
819      'traits': traits,
820      'description': description,
821      'arguments': args,
822      'results': results,
823      'extras': extras
824  }
825
826
827def update_td_op_definitions(path, instructions, docs, filter_list,
828                             inst_category, capability_mapping):
829  """Updates SPIRVOps.td with newly generated op definition.
830
831  Arguments:
832    - path: path to SPIRVOps.td
833    - instructions: SPIR-V JSON grammar for all instructions
834    - docs: SPIR-V HTML doc for all instructions
835    - filter_list: a list containing new opnames to include
836    - capability_mapping: mapping from duplicated capability symbols to the
837                   canonicalized symbol chosen for SPIRVBase.td.
838
839  Returns:
840    - A string containing all the TableGen op definitions
841  """
842  with open(path, 'r') as f:
843    content = f.read()
844
845  # Split the file into chunks, each containing one op.
846  ops = content.split(AUTOGEN_OP_DEF_SEPARATOR)
847  header = ops[0]
848  footer = ops[-1]
849  ops = ops[1:-1]
850
851  # For each existing op, extract the manually-written sections out to retain
852  # them when re-generating the ops. Also append the existing ops to filter
853  # list.
854  name_op_map = {}  # Map from opname to its existing ODS definition
855  op_info_dict = {}
856  for op in ops:
857    info_dict = extract_td_op_info(op)
858    opname = info_dict['opname']
859    name_op_map[opname] = op
860    op_info_dict[opname] = info_dict
861    filter_list.append(opname)
862  filter_list = sorted(list(set(filter_list)))
863
864  op_defs = []
865  for opname in filter_list:
866    # Find the grammar spec for this op
867    try:
868      instruction = next(
869          inst for inst in instructions if inst['opname'] == opname)
870      op_defs.append(
871          get_op_definition(
872              instruction, docs[opname],
873              op_info_dict.get(opname, {'inst_category': inst_category}),
874              capability_mapping))
875    except StopIteration:
876      # This is an op added by us; use the existing ODS definition.
877      op_defs.append(name_op_map[opname])
878
879  # Substitute the old op definitions
880  op_defs = [header] + op_defs + [footer]
881  content = AUTOGEN_OP_DEF_SEPARATOR.join(op_defs)
882
883  with open(path, 'w') as f:
884    f.write(content)
885
886
887if __name__ == '__main__':
888  import argparse
889
890  cli_parser = argparse.ArgumentParser(
891      description='Update SPIR-V dialect definitions using SPIR-V spec')
892
893  cli_parser.add_argument(
894      '--base-td-path',
895      dest='base_td_path',
896      type=str,
897      default=None,
898      help='Path to SPIRVBase.td')
899  cli_parser.add_argument(
900      '--op-td-path',
901      dest='op_td_path',
902      type=str,
903      default=None,
904      help='Path to SPIRVOps.td')
905
906  cli_parser.add_argument(
907      '--new-enum',
908      dest='new_enum',
909      type=str,
910      default=None,
911      help='SPIR-V enum to be added to SPIRVBase.td')
912  cli_parser.add_argument(
913      '--new-opcodes',
914      dest='new_opcodes',
915      type=str,
916      default=None,
917      nargs='*',
918      help='update SPIR-V opcodes in SPIRVBase.td')
919  cli_parser.add_argument(
920      '--new-inst',
921      dest='new_inst',
922      type=str,
923      default=None,
924      nargs='*',
925      help='SPIR-V instruction to be added to ops file')
926  cli_parser.add_argument(
927      '--inst-category',
928      dest='inst_category',
929      type=str,
930      default='Op',
931      help='SPIR-V instruction category used for choosing '\
932           'the TableGen base class to define this op')
933
934  args = cli_parser.parse_args()
935
936  operand_kinds, instructions = get_spirv_grammar_from_json_spec()
937
938  # Define new enum attr
939  if args.new_enum is not None:
940    assert args.base_td_path is not None
941    filter_list = [args.new_enum] if args.new_enum else []
942    update_td_enum_attrs(args.base_td_path, operand_kinds, filter_list)
943
944  # Define new opcode
945  if args.new_opcodes is not None:
946    assert args.base_td_path is not None
947    update_td_opcodes(args.base_td_path, instructions, args.new_opcodes)
948
949  # Define new op
950  if args.new_inst is not None:
951    assert args.op_td_path is not None
952    docs = get_spirv_doc_from_html_spec()
953    capability_mapping = get_capability_mapping(operand_kinds)
954    update_td_op_definitions(args.op_td_path, instructions, docs, args.new_inst,
955                             args.inst_category, capability_mapping)
956    print('Done. Note that this script just generates a template; ', end='')
957    print('please read the spec and update traits, arguments, and ', end='')
958    print('results accordingly.')
959