hyf

Context-aware query service for Radroots
git clone https://radroots.dev/git/hyf.git
Log | Files | Refs | README | LICENSE

query_analysis.mojo (11865B)


      1 from std.collections import List, Optional
      2 
      3 from json import Value, loads
      4 
      5 from hyf_core.provenance import (
      6     CoreResponseMeta,
      7     ExecutionProvenance,
      8     ProvenanceSourceRef,
      9 )
     10 from hyf_core.request_context import RequestContext
     11 
     12 
     13 def _require_object(value: Value, context: String) raises:
     14     if not value.is_object():
     15         raise Error(context + " must be a JSON object")
     16 
     17 
     18 def _require_allowed_keys(
     19     value: Value, key_a: String, key_b: String, context: String
     20 ) raises:
     21     for key in value.object_keys():
     22         if key != key_a and key != key_b:
     23             raise Error(context + " contains unexpected field '" + key + "'")
     24 
     25 
     26 def has_key(value: Value, key: String) -> Bool:
     27     for candidate in value.object_keys():
     28         if candidate == key:
     29             return True
     30     return False
     31 
     32 
     33 def copy_string_list(items: List[String]) -> List[String]:
     34     var copied = List[String]()
     35     for item in items:
     36         copied.append(String(item))
     37     return copied^
     38 
     39 
     40 def string_array_value(items: List[String]) raises -> Value:
     41     var array = loads("[]")
     42     for item in items:
     43         array.append(Value(String(item)))
     44     return array^
     45 
     46 
     47 def collapse_whitespace(text: String) -> String:
     48     var parts = text.split()
     49     var collapsed = String()
     50     var first = True
     51     for part in parts:
     52         if not first:
     53             collapsed += " "
     54         collapsed += String(part)
     55         first = False
     56     return collapsed^
     57 
     58 
     59 def join_strings(items: List[String]) -> String:
     60     var joined = String()
     61     var first = True
     62     for item in items:
     63         if not first:
     64             joined += " "
     65         joined += String(item)
     66         first = False
     67     return joined^
     68 
     69 
     70 def normalize_free_text(text: String, mut signals: List[String]) -> String:
     71     var normalized = text.lower()
     72     if normalized != text:
     73         signals.append("lowercase")
     74 
     75     var replaced = normalized
     76     replaced = replaced.replace(",", " ")
     77     replaced = replaced.replace(".", " ")
     78     replaced = replaced.replace("!", " ")
     79     replaced = replaced.replace("?", " ")
     80     replaced = replaced.replace(":", " ")
     81     replaced = replaced.replace(";", " ")
     82     replaced = replaced.replace("/", " ")
     83     replaced = replaced.replace("\\", " ")
     84     replaced = replaced.replace("(", " ")
     85     replaced = replaced.replace(")", " ")
     86     replaced = replaced.replace("[", " ")
     87     replaced = replaced.replace("]", " ")
     88     replaced = replaced.replace("{", " ")
     89     replaced = replaced.replace("}", " ")
     90     replaced = replaced.replace('"', " ")
     91     replaced = replaced.replace("'", " ")
     92     replaced = replaced.replace("-", " ")
     93     if replaced != normalized:
     94         signals.append("punctuation_trimmed")
     95 
     96     var collapsed = collapse_whitespace(replaced)
     97     if collapsed != replaced:
     98         signals.append("whitespace_collapsed")
     99 
    100     return collapsed^
    101 
    102 
    103 def contains_token(items: List[String], token: String) -> Bool:
    104     for item in items:
    105         if item == token:
    106             return True
    107     return False
    108 
    109 
    110 def _is_stop_word(token: String) -> Bool:
    111     return (
    112         token == "a"
    113         or token == "an"
    114         or token == "and"
    115         or token == "for"
    116         or token == "from"
    117         or token == "in"
    118         or token == "me"
    119         or token == "near"
    120         or token == "of"
    121         or token == "on"
    122         or token == "the"
    123         or token == "to"
    124         or token == "with"
    125     )
    126 
    127 
    128 @fieldwise_init
    129 struct ExtractedFilters(Copyable, Movable):
    130     var local_intent: Bool
    131     var fulfillment: String
    132     var time_window: String
    133 
    134 
    135 @fieldwise_init
    136 struct QueryAnalysis(Copyable, Movable):
    137     var original_text: String
    138     var normalized_text: String
    139     var rewritten_text: String
    140     var query_terms: List[String]
    141     var normalization_signals: List[String]
    142     var ranking_hints: List[String]
    143     var extracted_filters: ExtractedFilters
    144 
    145 
    146 @fieldwise_init
    147 struct QueryRewriteRequest(Copyable, Movable):
    148     var text: String
    149 
    150 
    151 def extract_text_input(input: Value, capability_name: String) raises -> String:
    152     if not input.is_object():
    153         raise Error(capability_name + " input must be a JSON object")
    154 
    155     if has_key(input, "text"):
    156         var text_value = input["text"]
    157         if not text_value.is_string():
    158             raise Error(
    159                 capability_name + " input field 'text' must be a string"
    160             )
    161         var collapsed = collapse_whitespace(text_value.string_value())
    162         if collapsed == "":
    163             raise Error(capability_name + " input text must not be empty")
    164         return collapsed^
    165     elif has_key(input, "query"):
    166         var query_value = input["query"]
    167         if not query_value.is_string():
    168             raise Error(
    169                 capability_name + " input field 'query' must be a string"
    170             )
    171         var collapsed = collapse_whitespace(query_value.string_value())
    172         if collapsed == "":
    173             raise Error(capability_name + " input text must not be empty")
    174         return collapsed^
    175     else:
    176         raise Error(capability_name + " input requires 'text' or 'query'")
    177 
    178 
    179 def parse_query_rewrite_request(input: Value) raises -> QueryRewriteRequest:
    180     _require_object(input, "query_rewrite input")
    181     _require_allowed_keys(input, "text", "query", "query_rewrite input")
    182 
    183     var has_text = has_key(input, "text")
    184     var has_query = has_key(input, "query")
    185 
    186     if has_text and has_query:
    187         raise Error(
    188             "query_rewrite input must provide exactly one of 'text' or 'query'"
    189         )
    190     if not has_text and not has_query:
    191         raise Error(
    192             "query_rewrite input requires exactly one of 'text' or 'query'"
    193         )
    194 
    195     var source_field = "text" if has_text else "query"
    196     var text_value = input[source_field]
    197     if not text_value.is_string():
    198         raise Error(
    199             "query_rewrite input field '" + source_field + "' must be a string"
    200         )
    201 
    202     var collapsed = collapse_whitespace(text_value.string_value())
    203     if collapsed == "":
    204         raise Error("query_rewrite input text must not be empty")
    205 
    206     return QueryRewriteRequest(text=collapsed)
    207 
    208 
    209 def analyze_query_text(
    210     original_text: String, context: RequestContext
    211 ) -> QueryAnalysis:
    212     var normalized_input = String(original_text)
    213 
    214     var normalization_signals = List[String]()
    215     var normalized_text = normalize_free_text(
    216         normalized_input, normalization_signals
    217     )
    218     var normalized_tokens = normalized_text.split()
    219 
    220     var query_terms = List[String]()
    221     var ranking_hints = List[String]()
    222     var local_intent = False
    223     var fulfillment = "unspecified"
    224     var time_window = "unspecified"
    225     var removed_stop_words = False
    226     var extracted_filter_tokens = False
    227 
    228     for raw_token in normalized_tokens:
    229         var token = String(raw_token)
    230         if token == "":
    231             continue
    232 
    233         if (
    234             token == "near"
    235             or token == "me"
    236             or token == "nearby"
    237             or token == "local"
    238         ):
    239             local_intent = True
    240             extracted_filter_tokens = True
    241             continue
    242 
    243         if token == "pickup" or token == "curbside":
    244             fulfillment = "pickup"
    245             extracted_filter_tokens = True
    246             continue
    247 
    248         if token == "delivery" or token == "ship" or token == "shipping":
    249             fulfillment = "delivery"
    250             extracted_filter_tokens = True
    251             continue
    252 
    253         if token == "weekend" or token == "saturday" or token == "sunday":
    254             time_window = "weekend"
    255             extracted_filter_tokens = True
    256             continue
    257 
    258         if _is_stop_word(token):
    259             removed_stop_words = True
    260             continue
    261 
    262         if not contains_token(query_terms, token):
    263             query_terms.append(token)
    264 
    265     if local_intent:
    266         normalization_signals.append("local_intent_detected")
    267         ranking_hints.append("prefer_local_results")
    268     if fulfillment == "pickup":
    269         normalization_signals.append("pickup_filter_detected")
    270         ranking_hints.append("prefer_pickup")
    271     elif fulfillment == "delivery":
    272         normalization_signals.append("delivery_filter_detected")
    273         ranking_hints.append("prefer_delivery")
    274     if time_window == "weekend":
    275         normalization_signals.append("weekend_filter_detected")
    276         ranking_hints.append("prefer_weekend_availability")
    277     if removed_stop_words:
    278         normalization_signals.append("stopwords_removed")
    279     if extracted_filter_tokens:
    280         normalization_signals.append("filter_tokens_extracted")
    281     if context.scope:
    282         ranking_hints.append("respect_scope")
    283         normalization_signals.append("scope_present")
    284 
    285     if len(query_terms) == 0:
    286         query_terms.append(String(normalized_text))
    287         normalization_signals.append("fallback_to_normalized_query")
    288 
    289     return QueryAnalysis(
    290         original_text=normalized_input,
    291         normalized_text=normalized_text,
    292         rewritten_text=join_strings(query_terms),
    293         query_terms=query_terms^,
    294         normalization_signals=normalization_signals^,
    295         ranking_hints=ranking_hints^,
    296         extracted_filters=ExtractedFilters(
    297             local_intent=local_intent,
    298             fulfillment=fulfillment,
    299             time_window=time_window,
    300         ),
    301     )
    302 
    303 
    304 def analyze_query(
    305     input: Value, context: RequestContext, capability_name: String
    306 ) raises -> QueryAnalysis:
    307     var original_text = extract_text_input(input, capability_name)
    308     return analyze_query_text(original_text, context)
    309 
    310 
    311 def serialize_extracted_filters(filters: ExtractedFilters) raises -> Value:
    312     var value = loads("{}")
    313     value.set("local_intent", Value(filters.local_intent))
    314     value.set("fulfillment", Value(String(filters.fulfillment)))
    315     value.set("time_window", Value(String(filters.time_window)))
    316     return value^
    317 
    318 
    319 def query_signal_tags(analysis: QueryAnalysis) -> List[String]:
    320     var signal_tags = copy_string_list(analysis.normalization_signals)
    321     for hint in analysis.ranking_hints:
    322         signal_tags.append(String(hint))
    323     return signal_tags^
    324 
    325 
    326 def build_deterministic_meta(
    327     context: RequestContext,
    328     capability_name: String,
    329     signal_tags: List[String],
    330     extra_source_refs: List[ProvenanceSourceRef],
    331 ) -> CoreResponseMeta:
    332     var source_refs = List[ProvenanceSourceRef]()
    333     source_refs.append(
    334         ProvenanceSourceRef(
    335             source_kind="local_input",
    336             source_ref=capability_name + ":input",
    337         )
    338     )
    339     for source_ref in extra_source_refs:
    340         source_refs.append(
    341             ProvenanceSourceRef(
    342                 source_kind=String(source_ref.source_kind),
    343                 source_ref=String(source_ref.source_ref),
    344             )
    345         )
    346     if context.scope:
    347         source_refs.append(
    348             ProvenanceSourceRef(
    349                 source_kind="request_scope",
    350                 source_ref="request_context.scope",
    351             )
    352         )
    353 
    354     if context.return_provenance:
    355         return CoreResponseMeta(
    356             execution_mode="deterministic",
    357             backend="heuristic",
    358             provider=None,
    359             route=None,
    360             model=None,
    361             latency_ms=None,
    362             schema_version=Optional[Int](1),
    363             prompt_version=None,
    364             fallback_kind=None,
    365             fallback_reason=None,
    366             provenance=ExecutionProvenance(
    367                 kind="deterministic",
    368                 signal_tags=copy_string_list(signal_tags),
    369                 source_refs=source_refs^,
    370                 fallback=None,
    371                 evidence_set_id=None,
    372             ),
    373         )
    374 
    375     return CoreResponseMeta(
    376         execution_mode="deterministic",
    377         backend="heuristic",
    378         provider=None,
    379         route=None,
    380         model=None,
    381         latency_ms=None,
    382         schema_version=Optional[Int](1),
    383         prompt_version=None,
    384         fallback_kind=None,
    385         fallback_reason=None,
    386         provenance=None,
    387     )