Skip to main content

meta_language/grammar/inference/advisor/
llm.rs

1use super::{
2    validate_name_candidate, AdviceSource, ConceptNamingAdvisor, MdlMergeAdvisor, MergeAdvisor,
3    MergeRequest, MergeScore, NameCandidate, NamingAdvisor, NamingRequest,
4    INFERENCE_NAMING_CONCEPTS,
5};
6
7/// Provider-agnostic LLM boundary for optional inference acceleration.
8pub trait LlmClient: Send + Sync {
9    /// Completes one prompt.
10    fn complete(&self, prompt: &str) -> Result<String, LlmError>;
11}
12
13/// Error returned by an optional LLM client.
14#[derive(Clone, Debug, PartialEq, Eq)]
15pub struct LlmError {
16    message: String,
17}
18
19impl LlmError {
20    /// Builds an LLM client error.
21    #[must_use]
22    pub fn new(message: impl Into<String>) -> Self {
23        Self {
24            message: message.into(),
25        }
26    }
27
28    /// Human-readable error message.
29    #[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/// Optional LLM-backed naming advisor with deterministic fallback.
44#[derive(Clone, Debug, PartialEq, Eq)]
45pub struct LlmNamingAdvisor<C> {
46    client: C,
47    fallback: ConceptNamingAdvisor,
48}
49
50impl<C> LlmNamingAdvisor<C> {
51    /// Builds an LLM naming advisor with the default deterministic fallback.
52    #[must_use]
53    pub const fn new(client: C) -> Self {
54        Self {
55            client,
56            fallback: ConceptNamingAdvisor,
57        }
58    }
59
60    /// Builds an LLM naming advisor with an explicit deterministic fallback.
61    #[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/// Optional LLM-backed merge advisor with deterministic fallback.
91#[derive(Clone, Debug, PartialEq, Eq)]
92pub struct LlmMergeAdvisor<C> {
93    client: C,
94    fallback: MdlMergeAdvisor,
95}
96
97impl<C> LlmMergeAdvisor<C> {
98    /// Builds an LLM merge advisor with the default deterministic fallback.
99    #[must_use]
100    pub const fn new(client: C) -> Self {
101        Self {
102            client,
103            fallback: MdlMergeAdvisor,
104        }
105    }
106
107    /// Builds an LLM merge advisor with an explicit deterministic fallback.
108    #[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}