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)