meta_language/grammar/inference/advisor/
llm.rs1use super::{
2 validate_name_candidate, AdviceSource, ConceptNamingAdvisor, MdlMergeAdvisor, MergeAdvisor,
3 MergeRequest, MergeScore, NameCandidate, NamingAdvisor, NamingRequest,
4 INFERENCE_NAMING_CONCEPTS,
5};
6
7pub trait LlmClient: Send + Sync {
9 fn complete(&self, prompt: &str) -> Result<String, LlmError>;
11}
12
13#[derive(Clone, Debug, PartialEq, Eq)]
15pub struct LlmError {
16 message: String,
17}
18
19impl LlmError {
20 #[must_use]
22 pub fn new(message: impl Into<String>) -> Self {
23 Self {
24 message: message.into(),
25 }
26 }
27
28 #[must_use]
30 pub fn message(&self) -> &str {
31 &self.message
32 }
33}
34
35impl std::fmt::Display for LlmError {
36 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37 formatter.write_str(&self.message)
38 }
39}
40
41impl std::error::Error for LlmError {}
42
43#[derive(Clone, Debug, PartialEq, Eq)]
45pub struct LlmNamingAdvisor<C> {
46 client: C,
47 fallback: ConceptNamingAdvisor,
48}
49
50impl<C> LlmNamingAdvisor<C> {
51 #[must_use]
53 pub const fn new(client: C) -> Self {
54 Self {
55 client,
56 fallback: ConceptNamingAdvisor,
57 }
58 }
59
60 #[must_use]
62 pub const fn with_fallback(client: C, fallback: ConceptNamingAdvisor) -> Self {
63 Self { client, fallback }
64 }
65}
66
67impl<C> NamingAdvisor for LlmNamingAdvisor<C>
68where
69 C: LlmClient,
70{
71 fn propose_names(&self, request: &NamingRequest<'_>) -> Vec<NameCandidate> {
72 let deterministic = self.fallback.propose_names(request);
73 let Ok(response) = self.client.complete(&naming_prompt(request)) else {
74 return deterministic;
75 };
76
77 let candidates = parse_name_candidates(&response)
78 .into_iter()
79 .filter(|candidate| validate_name_candidate(request, candidate))
80 .collect::<Vec<_>>();
81
82 if candidates.is_empty() {
83 deterministic
84 } else {
85 candidates
86 }
87 }
88}
89
90#[derive(Clone, Debug, PartialEq, Eq)]
92pub struct LlmMergeAdvisor<C> {
93 client: C,
94 fallback: MdlMergeAdvisor,
95}
96
97impl<C> LlmMergeAdvisor<C> {
98 #[must_use]
100 pub const fn new(client: C) -> Self {
101 Self {
102 client,
103 fallback: MdlMergeAdvisor,
104 }
105 }
106
107 #[must_use]
109 pub const fn with_fallback(client: C, fallback: MdlMergeAdvisor) -> Self {
110 Self { client, fallback }
111 }
112}
113
114impl<C> MergeAdvisor for LlmMergeAdvisor<C>
115where
116 C: LlmClient,
117{
118 fn rank_merges(&self, request: &MergeRequest<'_>) -> Vec<MergeScore> {
119 let deterministic = self.fallback.rank_merges(request);
120 let Ok(response) = self.client.complete(&merge_prompt(request)) else {
121 return deterministic;
122 };
123 let Some(parsed) = parse_merge_scores(&response, request.candidates.len()) else {
124 return deterministic;
125 };
126
127 parsed
128 .into_iter()
129 .zip(deterministic)
130 .map(|(score, deterministic)| {
131 let mut score = score.clamp(0.0, 1.0);
132 if deterministic.score < 0.5 {
133 score = score.min(deterministic.score);
134 }
135 MergeScore {
136 score,
137 source: AdviceSource::Llm,
138 }
139 })
140 .collect()
141 }
142}
143
144fn naming_prompt(request: &NamingRequest<'_>) -> String {
145 let concepts = INFERENCE_NAMING_CONCEPTS
146 .iter()
147 .map(|concept| format!("{}={}", concept.id, concept.name))
148 .collect::<Vec<_>>()
149 .join(", ");
150 format!(
151 "Suggest one grammar rule name as name|concept.\nExpression: {}\nSamples: {:?}\nConcepts: {concepts}",
152 request.rule_expr, request.sample_yields
153 )
154}
155
156fn merge_prompt(request: &MergeRequest<'_>) -> String {
157 let candidates = request
158 .candidates
159 .iter()
160 .map(|candidate| format!("{}<-{}", candidate.winner, candidate.loser))
161 .collect::<Vec<_>>()
162 .join(", ");
163 format!("Score merge candidates in order with numbers from 0 to 1: {candidates}")
164}
165
166fn parse_name_candidates(response: &str) -> Vec<NameCandidate> {
167 response.lines().filter_map(parse_name_candidate).collect()
168}
169
170fn parse_name_candidate(line: &str) -> Option<NameCandidate> {
171 let line = line.trim();
172 if line.is_empty() {
173 return None;
174 }
175
176 let (name, concept) = line
177 .split_once('|')
178 .map_or((line, None), |(name, concept)| {
179 let concept = concept.trim();
180 let concept = if concept.is_empty() || concept.eq_ignore_ascii_case("none") {
181 None
182 } else {
183 Some(concept.to_string())
184 };
185 (name, concept)
186 });
187 let name = name.trim().trim_matches(['"', '\'', '`']);
188 (!name.is_empty()).then(|| NameCandidate {
189 name: name.to_string(),
190 concept,
191 source: AdviceSource::Llm,
192 })
193}
194
195fn parse_merge_scores(response: &str, expected_len: usize) -> Option<Vec<f64>> {
196 let scores = response
197 .split(|character: char| character.is_ascii_whitespace() || matches!(character, ',' | ';'))
198 .filter_map(|token| token.trim().parse::<f64>().ok())
199 .collect::<Vec<_>>();
200
201 (scores.len() == expected_len && scores.iter().all(|score| score.is_finite())).then_some(scores)
202}