File size: 17,970 Bytes
1c57e18
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7383a8e
1c57e18
 
 
 
 
 
 
 
 
 
 
 
 
 
7383a8e
1c57e18
 
 
 
7383a8e
 
1c57e18
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7383a8e
 
 
1c57e18
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7383a8e
1c57e18
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
"""
EasyRead Controller.

Orchestrates the EasyRead pipeline (simplify -> validate -> revise -> icons)
over the services in ``src/services`` and the prompts in ``src/prompts``.

Provider selection
------------------
* LLM provider (``llm_provider``): which model drives simplification,
  validation and revision. Defaults to ``"llama"`` (Meta Llama 3 70B on AWS
  Bedrock). Alternatives: ``"gemini"``, ``"bedrock"``.
* Translation provider (``translation_provider``): how sentences are
  translated. Defaults to ``"nllb"`` (facebook/nllb-200). Alternatives:
  ``"google"`` (Google Translate) and ``"gemini"``.

Heavy or credential-bound dependencies (the icon generator, translators) are
loaded lazily so constructing the Controller stays cheap and side-effect free.

Usage:
    from controller import Controller

    controller = Controller()                       # llama + nllb (defaults)
    controller = Controller(llm_provider="gemini",
                            translation_provider="google")

    result = controller.simplify_text(markdown, target_language="sw")
"""

import json
import re
import sys
from logging import getLogger
from pathlib import Path
from uuid import uuid4

import yaml

# Make ``src`` importable whether the controller is run from the repo root or
# imported as a module.
BASE_DIR = Path(__file__).resolve().parent
if str(BASE_DIR) not in sys.path:
    sys.path.insert(0, str(BASE_DIR))

logger = getLogger(__name__)

PROMPTS_DIR = BASE_DIR / "src" / "prompts"
ICON_OUTPUT_PATH = BASE_DIR / "src" / "temp"

# Supported providers (used for validation / friendly errors).
LLM_PROVIDERS = ("hf", "llama", "gemini", "bedrock")
TRANSLATION_PROVIDERS = ("nllb", "google", "gemini")


def _strip_code_fences(text: str) -> str:
    """Remove markdown code fences (```json ... ```) that LLMs sometimes add."""
    match = re.search(r'```(?:json)?\s*([\s\S]*?)```', text)
    return match.group(1).strip() if match else text.strip()


class Controller:
    """Drive the EasyRead pipeline with pluggable LLM and translation backends."""

    def __init__(
        self,
        llm_provider: str = "hf",
        translation_provider: str = "nllb",
    ):
        """
        Args:
            llm_provider: One of ``LLM_PROVIDERS``. Default ``"hf"`` (Llama via
                Hugging Face Inference Providers).
            translation_provider: One of ``TRANSLATION_PROVIDERS``. Default
                ``"nllb"``.
        """
        self.llm_provider = llm_provider.lower()
        self.translation_provider = translation_provider.lower()

        if self.llm_provider not in LLM_PROVIDERS:
            raise ValueError(
                f"Unknown llm_provider '{llm_provider}'. Choose from: {', '.join(LLM_PROVIDERS)}"
            )
        if self.translation_provider not in TRANSLATION_PROVIDERS:
            raise ValueError(
                f"Unknown translation_provider '{translation_provider}'. "
                f"Choose from: {', '.join(TRANSLATION_PROVIDERS)}"
            )

        # Eagerly build the selected LLM (the pipeline's core dependency).
        self.llm = self._make_llm(self.llm_provider)

        # Prompts loaded from src/prompts/*.yaml.
        self.prompts = self._load_prompts()

        self.icon_output_path = ICON_OUTPUT_PATH

        # Lazily-initialised collaborators.
        self._translator = None
        self._gemini = None
        self._icon_generator = None
        self._global_symbols = None

    # ------------------------------------------------------------------ #
    # Provider factories                                                  #
    # ------------------------------------------------------------------ #
    @staticmethod
    def _make_llm(provider: str):
        """Instantiate the chosen LLM driver. Each exposes generate_text(prompt)."""
        if provider == "hf":
            from src.services.hf_inference import HFInferenceDriver
            return HFInferenceDriver()
        if provider == "llama":
            from src.services.llama_3_70b import LlamaDriver
            return LlamaDriver()
        if provider == "gemini":
            from src.services.gemini import GeminiDriver
            return GeminiDriver()
        if provider == "bedrock":
            from src.services.bedrock import BedrockDriver
            return BedrockDriver()
        raise ValueError(f"Unknown llm_provider '{provider}'")

    def _get_translator(self):
        """Lazily build the chosen translation backend (google or nllb)."""
        if self._translator is not None:
            return self._translator

        if self.translation_provider == "google":
            from src.services.google_translate import GoogleTranslateDriver
            self._translator = GoogleTranslateDriver()
        elif self.translation_provider == "nllb":
            try:
                from src.services.nllb_200 import NLLBDriver
            except ImportError as exc:
                raise RuntimeError(
                    "NLLB translation driver (src/services/nllb_200.py) is not "
                    "implemented yet. Provide the NLLBDriver script, or construct "
                    "the Controller with translation_provider='google' or 'gemini'."
                ) from exc
            self._translator = NLLBDriver()
        else:  # pragma: no cover - gemini handled in translate_text
            raise RuntimeError(f"No standalone translator for '{self.translation_provider}'")

        return self._translator

    def _get_gemini(self):
        """Lazily build a GeminiDriver (used for gemini-based translation)."""
        if self._gemini is None:
            from src.services.gemini import GeminiDriver
            # Reuse the LLM instance if it already is Gemini.
            self._gemini = self.llm if self.llm_provider == "gemini" else GeminiDriver()
        return self._gemini

    @property
    def icon_generator(self):
        """Lazily load the ETH LoRA icon generator (pulls torch/diffusers)."""
        if self._icon_generator is None:
            from src.services.eth_zurich import IconGenerator
            self._icon_generator = IconGenerator()
        return self._icon_generator

    @property
    def global_symbols(self):
        """Lazily build the Global Symbols API client."""
        if self._global_symbols is None:
            from src.services.global_symbols import GlobalSymbolsService
            self._global_symbols = GlobalSymbolsService()
        return self._global_symbols

    # ------------------------------------------------------------------ #
    # Prompt loading                                                      #
    # ------------------------------------------------------------------ #
    @staticmethod
    def _load_prompts() -> dict:
        """Load all prompt YAMLs from src/prompts keyed by filename stem."""
        prompts = {}
        for path in PROMPTS_DIR.glob("*.yaml"):
            with open(path) as f:
                prompts[path.stem] = yaml.safe_load(f) or {}
        logger.info("Loaded prompts: %s", ", ".join(sorted(prompts)))
        return prompts

    # ------------------------------------------------------------------ #
    # Translation                                                         #
    # ------------------------------------------------------------------ #
    def translate_text(self, text: str, target_language: str) -> str:
        """
        Translate text using the configured translation provider.

        Args:
            text: Text to translate.
            target_language: Target language (ISO code for nllb/google).

        Returns:
            The translated text (or the original if ``target_language`` is empty).
        """
        if not target_language:
            return text

        if self.translation_provider == "gemini":
            gemini = self._get_gemini()
            prompt = (
                f"Translate the following text to {target_language}. "
                "Output only the translation, with no extra commentary or quotes.\n\n"
                f"Text:\n{text}"
            )
            return gemini.generate_text(prompt).strip()

        # google / nllb expose translate(text, target_language).
        return self._get_translator().translate(text, target_language)

    # ------------------------------------------------------------------ #
    # Core pipeline                                                       #
    # ------------------------------------------------------------------ #
    def simplify_text(self, text: str, target_language: str = None) -> dict:
        """Convert markdown into Easy Read sentences, optionally translated."""
        template = self.prompts['simplify_text']['system_message']
        prompt = f"{template}\n\n Bellow is the Input Text to simplify:\n\n{text}\n\n"

        response = self.llm.generate_text(prompt)

        try:
            response_data = json.loads(_strip_code_fences(response))
            logger.info("Successfully parsed simplify response JSON")

            if target_language:
                for s in response_data.get('simplified_sentences', []):
                    s['translated_sentence'] = self.translate_text(s['sentence'], target_language)

        except json.JSONDecodeError:
            logger.error(f"Failed to parse response as JSON. Raw response: {response}")
            response_data = {"error": "Failed to parse response as JSON.", "raw_response": response}
        return response_data

    def validate_text(self, original_sentence: str, simplified_sentences: list[dict]) -> dict:
        """Validate Easy Read sentences against the original markdown."""
        template = self.prompts['validate_text']['system_message']
        prompt = template + "\n" + self.prompts['validate_text']["user_message_template"].format(
            original_markdown=original_sentence,
            simplified_sentences=json.dumps(simplified_sentences),
        )

        response = self.llm.generate_text(prompt)

        try:
            response_data = json.loads(_strip_code_fences(response))
            logger.info("Successfully parsed validate response JSON")
        except json.JSONDecodeError:
            response_data = {"error": "Failed to parse response as JSON.", "raw_response": response}
            logger.error(f"Failed to parse response as JSON. Raw response: {response}")
        return response_data

    def revise_text(
        self,
        original_text: str,
        easy_read_sentences: list,
        feedback: str,
        target_language: str = None,
    ) -> dict:
        """Revise Easy Read sentences using validation feedback, optionally translated."""
        template = self.prompts['revise_text']['system_message']
        prompt = template + "\n" + self.prompts['revise_text']["user_message_template"].format(
            original_markdown=original_text,
            simplified_sentences=json.dumps(easy_read_sentences),
            validation_feedback=feedback,
        )

        response = self.llm.generate_text(prompt)

        try:
            response_data = json.loads(_strip_code_fences(response))
            logger.info("Successfully parsed revise response JSON")

            if target_language:
                for s in response_data.get('revised_sentences', []):
                    s['translated_sentence'] = self.translate_text(s['sentence'], target_language)

        except json.JSONDecodeError:
            response_data = {"error": "Failed to parse response as JSON.", "raw_response": response}
            logger.error(f"Failed to parse response as JSON. Raw response: {response}")

        return response_data

    # ------------------------------------------------------------------ #
    # Icons / symbols                                                     #
    # ------------------------------------------------------------------ #
    def generate_icons(
        self,
        sentences: list[dict],
        symbolset: str = "arasaac",
        use_global_symbols: bool = True,
    ) -> dict:
        """
        Generate icons for sentences using Global Symbols API, falling back to
        the local LoRA generator when no symbol is found.
        """
        request_id = str(uuid4())
        request_dir = Path(self.icon_output_path) / request_id
        request_dir.mkdir(parents=True, exist_ok=True)

        for sentence in sentences:
            prompt = sentence['image_prompt']
            safe_prompt = "_".join(prompt.split())
            image_path = request_dir / f"{safe_prompt}.png"

            image_found = False

            if use_global_symbols:
                logger.info(f"Querying Global Symbols API for prompt: '{prompt}'")
                try:
                    downloaded_path = self.global_symbols.search_and_download(
                        query=prompt,
                        output_path=image_path,
                        symbolset=symbolset,
                    )
                    if downloaded_path:
                        logger.info(f"Retrieved symbol from Global Symbols for '{prompt}'")
                        image_found = True
                    else:
                        logger.warning(f"No symbol found in Global Symbols for '{prompt}'")
                except Exception as e:
                    logger.error(f"Error querying Global Symbols for '{prompt}': {e}")

            if not image_found:
                logger.info(f"Generating icon locally for prompt: '{prompt}'")
                try:
                    image = self.icon_generator.generate(prompt)
                    image.save(image_path)
                    logger.info(f"Saved locally generated icon for '{prompt}'")
                except Exception as e:
                    logger.error(f"Error generating icon locally for '{prompt}': {e}")
                    continue

            sentence['image_path'] = "/".join(["icons", request_id, f"{safe_prompt}.png"])

        return {"request_id": request_id, "icons": sentences}

    def search_symbols(self, sentences: list[dict], symbolset: str = "arasaac") -> dict:
        """Search Global Symbols for each sentence without AI fallback."""
        request_id = str(uuid4())
        request_dir = Path(self.icon_output_path) / request_id
        request_dir.mkdir(parents=True, exist_ok=True)

        results = []
        for sentence in sentences:
            prompt = sentence['image_prompt']
            safe_prompt = "_".join(prompt.split())
            image_path = request_dir / f"{safe_prompt}.png"

            symbol_found = False
            symbol_image_path = None

            try:
                downloaded_path = self.global_symbols.search_and_download(
                    query=prompt,
                    output_path=image_path,
                    symbolset=symbolset,
                )
                if downloaded_path:
                    symbol_found = True
                    symbol_image_path = "/".join(["icons", request_id, f"{safe_prompt}.png"])
                    logger.info(f"Found symbol for '{prompt}'")
                else:
                    logger.info(f"No symbol found for '{prompt}'")
            except Exception as e:
                logger.error(f"Error searching symbols for '{prompt}': {e}")

            results.append({
                "sentence": sentence['sentence'],
                "image_prompt": sentence['image_prompt'],
                "highlighted": sentence['highlighted'],
                "symbol_found": symbol_found,
                "symbol_image_path": symbol_image_path,
            })

        return {"request_id": request_id, "results": results}

    def generate_ai_icons(self, request_id: str, sentences: list[dict]) -> dict:
        """
        Generate AI icons for sentences, passing through symbol search results
        where available and generating the rest with the ETH LoRA model.
        """
        request_dir = Path(self.icon_output_path) / request_id
        request_dir.mkdir(parents=True, exist_ok=True)

        for sentence in sentences:
            symbol_image_path = sentence.get('symbol_image_path')

            if symbol_image_path:
                sentence['image_path'] = symbol_image_path
                logger.info(f"Using existing symbol for '{sentence['sentence']}'")
            else:
                ai_prompt = sentence['ai_prompt']
                safe_prompt = "_".join(ai_prompt.split())[:80]
                image_path = request_dir / f"ai_{safe_prompt}.png"

                try:
                    image = self.icon_generator.generate(ai_prompt)
                    image.save(image_path)
                    logger.info(f"Saved AI-generated icon for '{ai_prompt}'")
                except Exception as e:
                    logger.error(f"Error generating AI icon for '{ai_prompt}': {e}")
                    continue

                sentence['image_path'] = "/".join(["icons", request_id, f"ai_{safe_prompt}.png"])

            sentence['image_prompt'] = sentence.get('ai_prompt', sentence.get('sentence', ''))

        return {"request_id": request_id, "icons": sentences}


if __name__ == "__main__":
    import argparse

    parser = argparse.ArgumentParser(description="EasyRead controller — simplify markdown into Easy Read")
    parser.add_argument("text", type=str, help="Markdown/text to simplify")
    parser.add_argument("--llm", type=str, default="hf", choices=LLM_PROVIDERS, help="LLM provider")
    parser.add_argument("--translation", type=str, default="nllb", choices=TRANSLATION_PROVIDERS, help="Translation provider")
    parser.add_argument("--target-language", type=str, default=None, help="Optional target language code (e.g. sw)")

    args = parser.parse_args()

    controller = Controller(llm_provider=args.llm, translation_provider=args.translation)
    output = controller.simplify_text(args.text, target_language=args.target_language)
    print(json.dumps(output, ensure_ascii=False, indent=2))