Brunobkr/llama.cpp_AlgMor24_github
ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.
03.1k
1#!/usr/bin/env python32from __future__ import annotations3 4import argparse5import itertools6import json7import re8import sys9from typing import Any, List, Optional, Set, Tuple, Union10 11def _build_repetition(item_rule, min_items, max_items, separator_rule=None):12 13 if max_items == 0:14 return ""15 16 if min_items == 0 and max_items == 1:17 return f'{item_rule}?'18 19 if not separator_rule:20 if min_items == 1 and max_items is None:21 return f'{item_rule}+'22 elif min_items == 0 and max_items is None:23 return f'{item_rule}*'24 else:25 return f'{item_rule}{{{min_items},{max_items if max_items is not None else ""}}}'26 27 result = item_rule + ' ' + _build_repetition(f'({separator_rule} {item_rule})', min_items - 1 if min_items > 0 else 0, max_items - 1 if max_items is not None else None)28 return f'({result})?' if min_items == 0 else result29 30def _generate_min_max_int(min_value: Optional[int], max_value: Optional[int], out: list, decimals_left: int = 16, top_level: bool = True):31 def digit_range(from_char: str, to_char: str):32 out.append("[")33 if from_char == to_char:34 out.append(from_char)35 else:36 out.append(from_char)37 out.append("-")38 out.append(to_char)39 out.append("]")40 41 def more_digits(min_digits: int, max_digits: int):42 out.append("[0-9]")43 if min_digits == max_digits and min_digits == 1:44 return45 out.append("{")46 out.append(str(min_digits))47 if max_digits != min_digits:48 out.append(",")49 if max_digits != sys.maxsize:50 out.append(str(max_digits))51 out.append("}")52 53 def uniform_range(from_str: str, to_str: str):54 i = 055 while i < len(from_str) and from_str[i] == to_str[i]:56 i += 157 if i > 0:58 out.append("\"")59 out.append(from_str[:i])60 out.append("\"")61 if i < len(from_str):62 if i > 0:63 out.append(" ")64 sub_len = len(from_str) - i - 165 if sub_len > 0:66 from_sub = from_str[i+1:]67 to_sub = to_str[i+1:]68 sub_zeros = "0" * sub_len69 sub_nines = "9" * sub_len70 71 to_reached = False72 out.append("(")73 if from_sub == sub_zeros:74 digit_range(from_str[i], chr(ord(to_str[i]) - 1))75 out.append(" ")76 more_digits(sub_len, sub_len)77 else:78 out.append("[")79 out.append(from_str[i])80 out.append("] ")81 out.append("(")82 uniform_range(from_sub, sub_nines)83 out.append(")")84 if ord(from_str[i]) < ord(to_str[i]) - 1:85 out.append(" | ")86 if to_sub == sub_nines:87 digit_range(chr(ord(from_str[i]) + 1), to_str[i])88 to_reached = True89 else:90 digit_range(chr(ord(from_str[i]) + 1), chr(ord(to_str[i]) - 1))91 out.append(" ")92 more_digits(sub_len, sub_len)93 if not to_reached:94 out.append(" | ")95 digit_range(to_str[i], to_str[i])96 out.append(" ")97 uniform_range(sub_zeros, to_sub)98 out.append(")")99 else:100 out.append("[")101 out.append(from_str[i])102 out.append("-")103 out.append(to_str[i])104 out.append("]")105 106 if min_value is not None and max_value is not None:107 if min_value < 0 and max_value < 0:108 out.append("\"-\" (")109 _generate_min_max_int(-max_value, -min_value, out, decimals_left, top_level=True)110 out.append(")")111 return112 113 if min_value < 0:114 out.append("\"-\" (")115 _generate_min_max_int(0, -min_value, out, decimals_left, top_level=True)116 out.append(") | ")117 min_value = 0118 119 min_s = str(min_value)120 max_s = str(max_value)121 min_digits = len(min_s)122 max_digits = len(max_s)123 124 for digits in range(min_digits, max_digits):125 uniform_range(min_s, "9" * digits)126 min_s = "1" + "0" * digits127 out.append(" | ")128 uniform_range(min_s, max_s)129 return130 131 less_decimals = max(decimals_left - 1, 1)132 133 if min_value is not None:134 if min_value < 0:135 out.append("\"-\" (")136 _generate_min_max_int(None, -min_value, out, decimals_left, top_level=False)137 out.append(") | [0] | [1-9] ")138 more_digits(0, decimals_left - 1)139 elif min_value == 0:140 if top_level:141 out.append("[0] | [1-9] ")142 more_digits(0, less_decimals)143 else:144 more_digits(1, decimals_left)145 elif min_value <= 9:146 c = str(min_value)147 range_start = '1' if top_level else '0'148 if c > range_start:149 digit_range(range_start, chr(ord(c) - 1))150 out.append(" ")151 more_digits(1, less_decimals)152 out.append(" | ")153 digit_range(c, "9")154 out.append(" ")155 more_digits(0, less_decimals)156 else:157 min_s = str(min_value)158 length = len(min_s)159 c = min_s[0]160 161 if c > "1":162 digit_range("1" if top_level else "0", chr(ord(c) - 1))163 out.append(" ")164 more_digits(length, less_decimals)165 out.append(" | ")166 digit_range(c, c)167 out.append(" (")168 _generate_min_max_int(int(min_s[1:]), None, out, less_decimals, top_level=False)169 out.append(")")170 if c < "9":171 out.append(" | ")172 digit_range(chr(ord(c) + 1), "9")173 out.append(" ")174 more_digits(length - 1, less_decimals)175 return176 177 if max_value is not None:178 if max_value >= 0:179 if top_level:180 out.append("\"-\" [1-9] ")181 more_digits(0, less_decimals)182 out.append(" | ")183 _generate_min_max_int(0, max_value, out, decimals_left, top_level=True)184 else:185 out.append("\"-\" (")186 _generate_min_max_int(-max_value, None, out, decimals_left, top_level=False)187 out.append(")")188 return189 190 raise RuntimeError("At least one of min_value or max_value must be set")191 192class BuiltinRule:193 def __init__(self, content: str, deps: list | None = None):194 self.content = content195 self.deps = deps or []196 197# Constraining spaces to prevent model "running away".198SPACE_RULE = '| " " | "\\n"{1,2} [ \\t]{0,20}'199 200PRIMITIVE_RULES = {201 'boolean' : BuiltinRule('("true" | "false")', []),202 'decimal-part' : BuiltinRule('[0-9]{1,16}', []),203 'integral-part': BuiltinRule('[0] | [1-9] [0-9]{0,15}', []),204 'number' : BuiltinRule('("-"? integral-part) ("." decimal-part)? ([eE] [-+]? integral-part)?', ['integral-part', 'decimal-part']),205 'integer' : BuiltinRule('("-"? integral-part)', ['integral-part']),206 'value' : BuiltinRule('object | array | string | number | boolean | null', ['object', 'array', 'string', 'number', 'boolean', 'null']),207 'object' : BuiltinRule('"{" space ( string ":" space value ("," space string ":" space value)* )? space "}"', ['string', 'value']),208 'array' : BuiltinRule('"[" space ( value ("," space value)* )? space "]"', ['value']),209 'uuid' : BuiltinRule(r'"\"" [0-9a-fA-F]{8} "-" [0-9a-fA-F]{4} "-" [0-9a-fA-F]{4} "-" [0-9a-fA-F]{4} "-" [0-9a-fA-F]{12} "\""', []),210 'char' : BuiltinRule(r'[^"\\\x7F\x00-\x1F] | [\\] (["\\bfnrt] | "u" [0-9a-fA-F]{4})', []),211 'string' : BuiltinRule(r'"\"" char* "\""', ['char']),212 'null' : BuiltinRule('"null"', []),213}214 215# TODO: support "uri", "email" string formats216STRING_FORMAT_RULES = {217 'date' : BuiltinRule('[0-9]{4} "-" ( "0" [1-9] | "1" [0-2] ) "-" ( \"0\" [1-9] | [1-2] [0-9] | "3" [0-1] )', []),218 'time' : BuiltinRule('([01] [0-9] | "2" [0-3]) ":" [0-5] [0-9] ":" [0-5] [0-9] ( "." [0-9]{3} )? ( "Z" | ( "+" | "-" ) ( [01] [0-9] | "2" [0-3] ) ":" [0-5] [0-9] )', []),219 'date-time' : BuiltinRule('date "T" time', ['date', 'time']),220 'date-string' : BuiltinRule('"\\"" date "\\""', ['date']),221 'time-string' : BuiltinRule('"\\"" time "\\""', ['time']),222 'date-time-string': BuiltinRule('"\\"" date-time "\\""', ['date-time']),223}224 225DOTALL = '[\\U00000000-\\U0010FFFF]'226DOT = '[^\\x0A\\x0D]'227 228RESERVED_NAMES = set(["root", "dot", *PRIMITIVE_RULES.keys(), *STRING_FORMAT_RULES.keys()])229 230INVALID_RULE_CHARS_RE = re.compile(r'[^a-zA-Z0-9-]+')231GRAMMAR_LITERAL_ESCAPE_RE = re.compile(r'[\r\n"\\]')232GRAMMAR_RANGE_LITERAL_ESCAPE_RE = re.compile(r'[\r\n"\]\-\\]')233GRAMMAR_LITERAL_ESCAPES = {'\r': '\\r', '\n': '\\n', '"': '\\"', '-': '\\-', ']': '\\]', '\\': '\\\\'}234 235NON_LITERAL_SET = set('|.()[]{}*+?')236ESCAPED_IN_REGEXPS_BUT_NOT_IN_LITERALS = set('^$.[]()|{}*+?')237 238 239class SchemaConverter:240 def __init__(self, *, prop_order, allow_fetch, dotall, raw_pattern):241 self._prop_order = prop_order242 self._allow_fetch = allow_fetch243 self._dotall = dotall244 self._raw_pattern = raw_pattern245 self._rules = {246 'space': SPACE_RULE,247 }248 self._refs = {}249 self._refs_being_resolved = set()250 251 def _format_literal(self, literal):252 escaped = GRAMMAR_LITERAL_ESCAPE_RE.sub(253 lambda m: GRAMMAR_LITERAL_ESCAPES.get(m.group(0)) or m.group(0), literal254 )255 return f'"{escaped}"'256 257 def not_literal(self, literal: str, dotall: bool = True, maybe_escaped_underscores = False) -> str:258 '''259 not_literal('a') -> '[^a]'260 not_literal('abc') -> '([^a] | "a" ([^b] | "b" ([^c])?)?)?'261 '''262 assert len(literal) > 0, 'Empty literal not supported'263 def recurse(i: int):264 c = literal[i]265 if maybe_escaped_underscores and c == '_':266 yield f'[^{c}\\\\]'267 yield ' | '268 yield f'"\\\\"? "{c}"'269 else:270 yield f'[^{c}]'271 if i < len(literal) - 1:272 yield ' | '273 yield self._format_literal(c)274 yield ' ('275 yield from recurse(i + 1)276 yield ')?'277 278 return ''.join(('(', *recurse(0), ')'))279 280 def _not_strings(self, strings):281 class TrieNode:282 def __init__(self):283 self.children = {}284 self.is_end_of_string = False285 286 def insert(self, string):287 node = self288 for c in string:289 node = node.children.setdefault(c, TrieNode())290 node.is_end_of_string = True291 292 trie = TrieNode()293 for s in strings:294 trie.insert(s)295 296 char_rule = self._add_primitive('char', PRIMITIVE_RULES['char'])297 out = ['["] ( ']298 299 def visit(node):300 rejects = []301 first = True302 for c in sorted(node.children.keys()):303 child = node.children[c]304 rejects.append(c)305 if first:306 first = False307 else:308 out.append(' | ')309 out.append(f'[{c}]')310 if child.children:311 out.append(f' (')312 visit(child)313 out.append(')')314 elif child.is_end_of_string:315 out.append(f' {char_rule}+')316 if node.children:317 if not first:318 out.append(' | ')319 out.append(f'[^"{"".join(rejects)}] {char_rule}*')320 visit(trie)321 322 out.append(f' ){"" if trie.is_end_of_string else "?"} ["]')323 return ''.join(out)324 325 def _add_rule(self, name, rule):326 esc_name = INVALID_RULE_CHARS_RE.sub('-', name)327 if esc_name not in self._rules or self._rules[esc_name] == rule:328 key = esc_name329 else:330 i = 0331 while f'{esc_name}{i}' in self._rules and self._rules[f'{esc_name}{i}'] != rule:332 i += 1333 key = f'{esc_name}{i}'334 self._rules[key] = rule335 return key336 337 def resolve_refs(self, schema: dict, url: str):338 '''339 Resolves all $ref fields in the given schema, fetching any remote schemas,340 replacing $ref with absolute reference URL and populating self._refs with the341 respective referenced (sub)schema dictionaries.342 '''343 def visit(n: dict):344 if isinstance(n, list):345 return [visit(x) for x in n]346 elif isinstance(n, dict):347 ref = n.get('$ref')348 if ref is not None and ref not in self._refs:349 if ref.startswith('https://'):350 assert self._allow_fetch, 'Fetching remote schemas is not allowed (use --allow-fetch for force)'351 import requests352 353 frag_split = ref.split('#')354 base_url = frag_split[0]355 356 target = self._refs.get(base_url)357 if target is None:358 target = self.resolve_refs(requests.get(ref).json(), base_url)359 self._refs[base_url] = target360 361 if len(frag_split) == 1 or frag_split[-1] == '':362 return target363 elif ref.startswith('#/'):364 target = schema365 ref = f'{url}{ref}'366 n['$ref'] = ref367 else:368 raise ValueError(f'Unsupported ref {ref}')369 370 for sel in ref.split('#')[-1].split('/')[1:]:371 assert target is not None, f'Error resolving ref {ref}: {sel} not in {target}'372 if isinstance(target, list):373 try:374 sel_index = int(sel)375 except ValueError:376 raise ValueError(f'Error resolving ref {ref}: {sel} not in {target}')377 assert 0 <= sel_index < len(target), f'Error resolving ref {ref}: {sel} not in {target}'378 target = target[sel_index]379 else:380 assert sel in target, f'Error resolving ref {ref}: {sel} not in {target}'381 target = target[sel]382 383 self._refs[ref] = target384 else:385 for v in n.values():386 visit(v)387 388 return n389 return visit(schema)390 391 def _generate_union_rule(self, name, alt_schemas):392 return ' | '.join((393 self.visit(alt_schema, f'{name}{"-" if name else "alternative-"}{i}')394 for i, alt_schema in enumerate(alt_schemas)395 ))396 397 def _visit_pattern(self, pattern, name):398 '''399 Transforms a regular expression pattern into a GBNF rule.400 401 Input: https://json-schema.org/understanding-json-schema/reference/regular_expressions402 Output: https://github.com/ggml-org/llama.cpp/blob/master/grammars/README.md403 404 Unsupported features: negative/positive lookaheads, greedy/non-greedy modifiers.405 406 Mostly a 1:1 translation, except for {x} / {x,} / {x,y} quantifiers for which407 we define sub-rules to keep the output lean.408 '''409 410 assert pattern.startswith('^') and pattern.endswith('$'), 'Pattern must start with "^" and end with "$"'411 pattern = pattern[1:-1]412 sub_rule_ids = {}413 414 i = 0415 length = len(pattern)416 417 def to_rule(s: tuple[str, bool]) -> str:418 (txt, is_literal) = s419 return "\"" + txt + "\"" if is_literal else txt420 421 def transform() -> tuple[str, bool]:422 '''423 Parse a unit at index i (advancing it), and return its string representation + whether it's a literal.424 '''425 nonlocal i426 nonlocal pattern427 nonlocal sub_rule_ids428 429 start = i430 # For each component of this sequence, store its string representation and whether it's a literal.431 # We only need a flat structure here to apply repetition operators to the last item, and432 # to merge literals at the and (we're parsing grouped ( sequences ) recursively and don't treat '|' specially433 # (GBNF's syntax is luckily very close to regular expressions!)434 seq: list[tuple[str, bool]] = []435 436 def get_dot():437 if self._dotall:438 rule = DOTALL439 else:440 # Accept any character... except \n and \r line break chars (\x0A and \xOD)441 rule = DOT442 return self._add_rule(f'dot', rule)443 444 def join_seq():445 nonlocal seq446 ret = []447 for is_literal, g in itertools.groupby(seq, lambda x: x[1]):448 if is_literal:449 ret.append((''.join(x[0] for x in g), True))450 else:451 ret.extend(g)452 if len(ret) == 1:453 return ret[0]454 return (' '.join(to_rule(x) for x in seq), False)455 456 while i < length:457 c = pattern[i]458 if c == '.':459 seq.append((get_dot(), False))460 i += 1461 elif c == '(':462 i += 1463 if i < length:464 assert pattern[i] != '?', f'Unsupported pattern syntax "{pattern[i]}" at index {i} of /{pattern}/'465 seq.append((f'({to_rule(transform())})', False))466 elif c == ')':467 i += 1468 assert start > 0 and pattern[start-1] == '(', f'Unbalanced parentheses; start = {start}, i = {i}, pattern = {pattern}'469 return join_seq()470 elif c == '[':471 square_brackets = c472 i += 1473 while i < length and pattern[i] != ']':474 if pattern[i] == '\\':475 square_brackets += pattern[i:i+2]476 i += 2477 else:478 square_brackets += pattern[i]479 i += 1480 assert i < length, f'Unbalanced square brackets; start = {start}, i = {i}, pattern = {pattern}'481 square_brackets += ']'482 i += 1483 seq.append((square_brackets, False))484 elif c == '|':485 seq.append(('|', False))486 i += 1487 elif c in ('*', '+', '?'):488 seq[-1] = (to_rule(seq[-1]) + c, False)489 i += 1490 elif c == '{':491 curly_brackets = c492 i += 1493 while i < length and pattern[i] != '}':494 curly_brackets += pattern[i]495 i += 1496 assert i < length, f'Unbalanced curly brackets; start = {start}, i = {i}, pattern = {pattern}'497 curly_brackets += '}'498 i += 1499 nums = [s.strip() for s in curly_brackets[1:-1].split(',')]500 min_times = 0501 max_times = None502 try:503 if len(nums) == 1:504 min_times = int(nums[0])505 max_times = min_times506 else:507 assert len(nums) == 2508 min_times = int(nums[0]) if nums[0] else 0509 max_times = int(nums[1]) if nums[1] else None510 except ValueError:511 raise ValueError(f'Invalid quantifier {curly_brackets} in /{pattern}/')512 513 (sub, sub_is_literal) = seq[-1]514 515 if not sub_is_literal:516 id = sub_rule_ids.get(sub)517 if id is None:518 id = self._add_rule(f'{name}-{len(sub_rule_ids) + 1}', sub)519 sub_rule_ids[sub] = id520 sub = id521 522 seq[-1] = (_build_repetition(f'"{sub}"' if sub_is_literal else sub, min_times, max_times), False)523 else:524 literal = ''525 while i < length:526 if pattern[i] == '\\' and i < length - 1:527 next = pattern[i + 1]528 if next in ESCAPED_IN_REGEXPS_BUT_NOT_IN_LITERALS:529 i += 1530 literal += pattern[i]531 i += 1532 else:533 literal += pattern[i:i+2]534 i += 2535 elif pattern[i] == '"' and not self._raw_pattern:536 literal += '\\"'537 i += 1538 elif pattern[i] not in NON_LITERAL_SET and \539 (i == length - 1 or literal == '' or pattern[i+1] == '.' or pattern[i+1] not in NON_LITERAL_SET):540 literal += pattern[i]541 i += 1542 else:543 break544 if literal:545 seq.append((literal, True))546 547 return join_seq()548 549 return self._add_rule(550 name,551 to_rule(transform()) if self._raw_pattern \552 else "\"\\\"\" (" + to_rule(transform()) + ") \"\\\"\"")553 554 555 def _resolve_ref(self, ref):556 ref_fragment = ref.split('#')[-1]557 ref_name = 'ref' + re.sub(r'[^a-zA-Z0-9-]+', '-', ref_fragment)558 if ref_name not in self._rules and ref not in self._refs_being_resolved:559 self._refs_being_resolved.add(ref)560 resolved = self._refs[ref]561 ref_name = self.visit(resolved, ref_name)562 self._refs_being_resolved.remove(ref)563 return ref_name564 565 def _generate_constant_rule(self, value):566 return self._format_literal(json.dumps(value))567 568 def visit(self, schema, name):569 schema_type = schema.get('type')570 schema_format = schema.get('format')571 rule_name = name + '-' if name in RESERVED_NAMES else name or 'root'572 573 if (ref := schema.get('$ref')) is not None:574 return self._add_rule(rule_name, self._resolve_ref(ref))575 576 elif 'oneOf' in schema or 'anyOf' in schema:577 return self._add_rule(rule_name, self._generate_union_rule(name, schema.get('oneOf') or schema['anyOf']))578 579 elif isinstance(schema_type, list):580 return self._add_rule(rule_name, self._generate_union_rule(name, [{**schema, 'type': t} for t in schema_type]))581 582 elif 'const' in schema:583 return self._add_rule(rule_name, self._generate_constant_rule(schema['const']))584 585 elif 'enum' in schema:586 rule = '(' + ' | '.join((self._generate_constant_rule(v) for v in schema['enum'])) + ')'587 return self._add_rule(rule_name, rule)588 589 elif schema_type in (None, 'object') and \590 ('properties' in schema or \591 ('additionalProperties' in schema and schema['additionalProperties'] is not True)):592 required = set(schema.get('required', []))593 properties = list(schema.get('properties', {}).items())594 return self._add_rule(rule_name, self._build_object_rule(properties, required, name, schema.get('additionalProperties')))595 596 elif schema_type in (None, 'object', 'string') and 'allOf' in schema:597 required = set()598 properties = []599 enum_sets = []600 hybrid_name = name601 def add_component(comp_schema, is_required):602 if (ref := comp_schema.get('$ref')) is not None:603 comp_schema = self._refs[ref]604 605 if 'properties' in comp_schema:606 for prop_name, prop_schema in comp_schema['properties'].items():607 properties.append((prop_name, prop_schema))608 if is_required:609 required.add(prop_name)610 611 if 'enum' in comp_schema:612 enum_sets.append(set(comp_schema['enum']))613 614 for t in schema['allOf']:615 if 'anyOf' in t:616 for tt in t['anyOf']:617 add_component(tt, is_required=False)618 else:619 add_component(t, is_required=True)620 621 if enum_sets:622 enum_intersection = enum_sets[0]623 for s in enum_sets[1:]:624 enum_intersection &= s625 626 if enum_intersection:627 rule = '(' + ' | '.join((self._generate_constant_rule(v) for v in sorted(enum_intersection))) + ')'628 return self._add_rule(rule_name, rule)629 630 return self._add_rule(rule_name, self._build_object_rule(properties, required, hybrid_name, additional_properties=None))631 632 elif schema_type in (None, 'array') and ('items' in schema or 'prefixItems' in schema):633 items = schema.get('items', schema.get('prefixItems'))634 if isinstance(items, list):635 return self._add_rule(636 rule_name,637 '"[" space ' +638 ' "," space '.join(639 self.visit(item, f'{name}{"-" if name else ""}tuple-{i}')640 for i, item in enumerate(items)) +641 ' space "]"')642 else:643 item_rule_name = self.visit(items, f'{name}{"-" if name else ""}item')644 min_items = schema.get("minItems", 0)645 max_items = schema.get("maxItems")646 return self._add_rule(rule_name, '"[" space ' + _build_repetition(item_rule_name, min_items, max_items, separator_rule='"," space') + ' space "]"')647 648 elif schema_type in (None, 'string') and 'pattern' in schema:649 return self._visit_pattern(schema['pattern'], rule_name)650 651 elif schema_type in (None, 'string') and re.match(r'^uuid[1-5]?$', schema_format or ''):652 return self._add_primitive(653 'root' if rule_name == 'root' else schema_format,654 PRIMITIVE_RULES['uuid']655 )656 657 elif schema_type in (None, 'string') and f'{schema_format}-string' in STRING_FORMAT_RULES:658 prim_name = f'{schema_format}-string'659 return self._add_rule(rule_name, self._add_primitive(prim_name, STRING_FORMAT_RULES[prim_name]))660 661 elif schema_type == 'string' and ('minLength' in schema or 'maxLength' in schema):662 char_rule = self._add_primitive('char', PRIMITIVE_RULES['char'])663 min_len = schema.get('minLength', 0)664 max_len = schema.get('maxLength')665 666 return self._add_rule(rule_name, r'"\"" ' + _build_repetition(char_rule, min_len, max_len) + r' "\""')667 668 elif schema_type in (None, 'integer') and \669 ('minimum' in schema or 'exclusiveMinimum' in schema or 'maximum' in schema or 'exclusiveMaximum' in schema):670 min_value = None671 max_value = None672 if 'minimum' in schema:673 min_value = schema['minimum']674 elif 'exclusiveMinimum' in schema:675 min_value = schema['exclusiveMinimum'] + 1676 if 'maximum' in schema:677 max_value = schema['maximum']678 elif 'exclusiveMaximum' in schema:679 max_value = schema['exclusiveMaximum'] - 1680 681 out = ["("]682 _generate_min_max_int(min_value, max_value, out)683 out.append(")")684 return self._add_rule(rule_name, ''.join(out))685 686 elif (schema_type == 'object') or (len(schema) == 0):687 return self._add_rule(rule_name, self._add_primitive('object', PRIMITIVE_RULES['object']))688 689 elif schema_type is None and isinstance(schema, dict):690 # No type constraint and no recognized structural keywords (e.g. {"description": "..."}).691 # Per JSON Schema semantics this is equivalent to {} and accepts any value.692 return self._add_rule(rule_name, self._add_primitive('value', PRIMITIVE_RULES['value']))693 694 else:695 assert schema_type in PRIMITIVE_RULES, f'Unrecognized schema: {schema}'696 # TODO: support minimum, maximum, exclusiveMinimum, exclusiveMaximum at least for zero697 return self._add_primitive('root' if rule_name == 'root' else schema_type, PRIMITIVE_RULES[schema_type])698 699 def _add_primitive(self, name: str, rule: BuiltinRule):700 n = self._add_rule(name, rule.content)701 702 for dep in rule.deps:703 dep_rule = PRIMITIVE_RULES.get(dep) or STRING_FORMAT_RULES.get(dep)704 assert dep_rule, f'Rule {dep} not known'705 if dep not in self._rules:706 self._add_primitive(dep, dep_rule)707 return n708 709 def _build_object_rule(self, properties: List[Tuple[str, Any]], required: Set[str], name: str, additional_properties: Optional[Union[bool, Any]]):710 prop_order = self._prop_order711 # sort by position in prop_order (if specified) then by original order712 sorted_props = [kv[0] for _, kv in sorted(enumerate(properties), key=lambda ikv: (prop_order.get(ikv[1][0], len(prop_order)), ikv[0]))]713 714 prop_kv_rule_names = {}715 for prop_name, prop_schema in properties:716 prop_rule_name = self.visit(prop_schema, f'{name}{"-" if name else ""}{prop_name}')717 prop_kv_rule_names[prop_name] = self._add_rule(718 f'{name}{"-" if name else ""}{prop_name}-kv',719 fr'{self._format_literal(json.dumps(prop_name))} space ":" space {prop_rule_name}'720 )721 required_props = [k for k in sorted_props if k in required]722 optional_props = [k for k in sorted_props if k not in required]723 724 if additional_properties is not None and additional_properties != False:725 sub_name = f'{name}{"-" if name else ""}additional'726 value_rule = self.visit(additional_properties, f'{sub_name}-value') if isinstance(additional_properties, dict) else \727 self._add_primitive('value', PRIMITIVE_RULES['value'])728 key_rule = self._add_primitive('string', PRIMITIVE_RULES['string']) if not sorted_props \729 else self._add_rule(f'{sub_name}-k', self._not_strings(sorted_props))730 731 prop_kv_rule_names["*"] = self._add_rule(732 f'{sub_name}-kv',733 f'{key_rule} ":" space {value_rule}'734 )735 optional_props.append("*")736 737 rule = '"{" space '738 rule += ' "," space '.join(prop_kv_rule_names[k] for k in required_props)739 740 if optional_props:741 rule += ' ('742 if required_props:743 rule += ' "," space ( '744 745 def get_recursive_refs(ks, first_is_optional):746 [k, *rest] = ks747 kv_rule_name = prop_kv_rule_names[k]748 comma_ref = f'( "," space {kv_rule_name} )'749 if first_is_optional:750 res = comma_ref + ('*' if k == '*' else '?')751 else:752 res = kv_rule_name + (' ' + comma_ref + "*" if k == '*' else '')753 if len(rest) > 0:754 res += ' ' + self._add_rule(755 f'{name}{"-" if name else ""}{k}-rest',756 get_recursive_refs(rest, first_is_optional=True)757 )758 return res759 760 rule += ' | '.join(761 get_recursive_refs(optional_props[i:], first_is_optional=False)762 for i in range(len(optional_props))763 )764 if required_props:765 rule += ' )'766 rule += ' )?'767 768 rule += ' space "}"'769 770 return rule771 772 def format_grammar(self):773 return '\n'.join(774 f'{name} ::= {rule}'775 for name, rule in sorted(self._rules.items(), key=lambda kv: kv[0])776 )777 778 779def main(args_in = None):780 parser = argparse.ArgumentParser(781 description='''782 Generates a grammar (suitable for use in ./llama-cli) that produces JSON conforming to a783 given JSON schema. Only a subset of JSON schema features are supported; more may be784 added in the future.785 ''',786 )787 parser.add_argument(788 '--prop-order',789 default=[],790 type=lambda s: s.split(','),791 help='''792 comma-separated property names defining the order of precedence for object properties;793 properties not specified here are given lower precedence than those that are, and794 are kept in their original order from the schema. Required properties are always795 given precedence over optional properties.796 '''797 )798 parser.add_argument(799 '--allow-fetch',800 action='store_true',801 default=False,802 help='Whether to allow fetching referenced schemas over HTTPS')803 parser.add_argument(804 '--dotall',805 action='store_true',806 default=False,807 help='Whether to treat dot (".") as matching all chars including line breaks in regular expression patterns')808 parser.add_argument(809 '--raw-pattern',810 action='store_true',811 default=False,812 help='Treats string patterns as raw patterns w/o quotes (or quote escapes)')813 814 parser.add_argument('schema', help='file containing JSON schema ("-" for stdin)')815 args = parser.parse_args(args_in)816 817 if args.schema.startswith('https://'):818 url = args.schema819 import requests820 schema = requests.get(url).json()821 elif args.schema == '-':822 url = 'stdin'823 schema = json.load(sys.stdin)824 else:825 url = f'file://{args.schema}'826 with open(args.schema) as f:827 schema = json.load(f)828 converter = SchemaConverter(829 prop_order={name: idx for idx, name in enumerate(args.prop_order)},830 allow_fetch=args.allow_fetch,831 dotall=args.dotall,832 raw_pattern=args.raw_pattern)833 schema = converter.resolve_refs(schema, url)834 converter.visit(schema, '')835 print(converter.format_grammar())836 837 838if __name__ == '__main__':839 main()840 