hyf

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

max_local.mojo (6288B)


      1 from std.collections import Optional
      2 
      3 from json import Value, loads
      4 
      5 from hyf_assist.contract import max_local_query_rewrite_route
      6 from hyf_core.capabilities.query_analysis import QueryAnalysis
      7 from hyf_core.request_context import RequestContext
      8 from hyf_provider.client import post_max_local_chat_completion
      9 from hyf_provider.config import MaxLocalProviderConfig
     10 from hyf_provider.health import resolve_max_local_provider_status
     11 from hyf_provider.result import (
     12     MaxLocalProviderStatus,
     13     parse_query_analysis_from_chat_completion,
     14 )
     15 from hyf_provider.schema import (
     16     build_query_rewrite_request_body,
     17     query_rewrite_prompt_version,
     18     query_rewrite_schema_version,
     19 )
     20 
     21 
     22 @fieldwise_init
     23 struct MaxLocalQueryRewriteResult(Copyable, Movable):
     24     var analysis: QueryAnalysis
     25     var provider: String
     26     var route: String
     27     var model: String
     28     var latency_ms: Int
     29     var schema_version: Int
     30     var prompt_version: String
     31 
     32 
     33 @fieldwise_init
     34 struct MaxLocalQueryRewriteFailure(Copyable, Movable):
     35     var kind: String
     36     var reason: String
     37 
     38 
     39 @fieldwise_init
     40 struct MaxLocalQueryRewriteOutcome(Copyable, Movable):
     41     var result: Optional[MaxLocalQueryRewriteResult]
     42     var failure: Optional[MaxLocalQueryRewriteFailure]
     43 
     44 
     45 def _query_rewrite_success_outcome(
     46     result: MaxLocalQueryRewriteResult,
     47 ) -> MaxLocalQueryRewriteOutcome:
     48     return MaxLocalQueryRewriteOutcome(
     49         result=Optional[MaxLocalQueryRewriteResult](result.copy()),
     50         failure=Optional[MaxLocalQueryRewriteFailure](None),
     51     )
     52 
     53 
     54 def _query_rewrite_failure_outcome(
     55     kind: String, reason: String
     56 ) -> MaxLocalQueryRewriteOutcome:
     57     return MaxLocalQueryRewriteOutcome(
     58         result=Optional[MaxLocalQueryRewriteResult](None),
     59         failure=Optional[MaxLocalQueryRewriteFailure](
     60             MaxLocalQueryRewriteFailure(
     61                 kind=String(kind), reason=String(reason)
     62             )
     63         ),
     64     )
     65 
     66 
     67 def max_local_query_rewrite_failure_from_reason(
     68     reason: String,
     69 ) -> MaxLocalQueryRewriteFailure:
     70     if reason == "invalid_url":
     71         return MaxLocalQueryRewriteFailure(
     72             kind="transport", reason="invalid_url"
     73         )
     74     if reason == "timeout":
     75         return MaxLocalQueryRewriteFailure(kind="transport", reason="timeout")
     76     if reason == "connection_failed":
     77         return MaxLocalQueryRewriteFailure(
     78             kind="transport", reason="connection_failed"
     79         )
     80     if reason == "unknown_transport":
     81         return MaxLocalQueryRewriteFailure(
     82             kind="provider", reason="provider_error"
     83         )
     84     if reason == "provider_non_2xx":
     85         return MaxLocalQueryRewriteFailure(
     86             kind="http_status", reason="provider_non_2xx"
     87         )
     88     if reason == "provider_error_payload":
     89         return MaxLocalQueryRewriteFailure(
     90             kind="provider_payload", reason="provider_error_payload"
     91         )
     92     if reason == "provider_invalid_json":
     93         return MaxLocalQueryRewriteFailure(
     94             kind="provider_payload", reason="provider_invalid_json"
     95         )
     96     if reason == "provider_schema_invalid":
     97         return MaxLocalQueryRewriteFailure(
     98             kind="provider_payload", reason="provider_schema_invalid"
     99         )
    100     if reason == "provider_empty_choices":
    101         return MaxLocalQueryRewriteFailure(
    102             kind="provider_payload", reason="provider_empty_choices"
    103         )
    104     if reason == "provider_missing_content":
    105         return MaxLocalQueryRewriteFailure(
    106             kind="provider_payload", reason="provider_missing_content"
    107         )
    108     return MaxLocalQueryRewriteFailure(kind="provider", reason="provider_error")
    109 
    110 
    111 def _load_chat_completion_response_json(text: String) raises -> Value:
    112     try:
    113         return loads(text)
    114     except:
    115         raise Error("provider_invalid_json")
    116 
    117 
    118 def _parse_query_analysis_from_body(text: String) raises -> QueryAnalysis:
    119     return parse_query_analysis_from_chat_completion(
    120         _load_chat_completion_response_json(text)
    121     )
    122 
    123 
    124 def execute_query_rewrite_via_max_local_provider(
    125     config: MaxLocalProviderConfig, text: String, context: RequestContext
    126 ) raises -> MaxLocalQueryRewriteResult:
    127     var outcome = try_execute_query_rewrite_via_max_local_provider(
    128         config, text, context
    129     )
    130     if outcome.result:
    131         return outcome.result.value().copy()
    132     if outcome.failure:
    133         raise Error(String(outcome.failure.value().reason))
    134     raise Error("provider_error")
    135 
    136 
    137 def try_execute_query_rewrite_via_max_local_provider(
    138     config: MaxLocalProviderConfig, text: String, context: RequestContext
    139 ) -> MaxLocalQueryRewriteOutcome:
    140     var request_body: Value
    141     try:
    142         request_body = build_query_rewrite_request_body(config, text, context)
    143     except:
    144         return _query_rewrite_failure_outcome("provider", "provider_error")
    145 
    146     var transport = post_max_local_chat_completion(
    147         config,
    148         request_body^,
    149     )
    150     if transport.failure:
    151         var failure = max_local_query_rewrite_failure_from_reason(
    152             transport.failure.value().reason
    153         )
    154         return _query_rewrite_failure_outcome(
    155             String(failure.kind), String(failure.reason)
    156         )
    157 
    158     if transport.response:
    159         try:
    160             var response = transport.response.value().copy()
    161             var analysis = _parse_query_analysis_from_body(response.body_text)
    162             return _query_rewrite_success_outcome(
    163                 MaxLocalQueryRewriteResult(
    164                     analysis=analysis^,
    165                     provider="max_local",
    166                     route=max_local_query_rewrite_route(),
    167                     model=String(config.model),
    168                     latency_ms=response.latency_ms,
    169                     schema_version=query_rewrite_schema_version(),
    170                     prompt_version=query_rewrite_prompt_version(),
    171                 )
    172             )
    173         except e:
    174             var failure = max_local_query_rewrite_failure_from_reason(String(e))
    175             return _query_rewrite_failure_outcome(
    176                 String(failure.kind), String(failure.reason)
    177             )
    178 
    179     return _query_rewrite_failure_outcome("provider", "provider_error")
    180 
    181 
    182 def max_local_provider_status(
    183     config: MaxLocalProviderConfig,
    184 ) -> MaxLocalProviderStatus:
    185     return resolve_max_local_provider_status(config)