File size: 42,254 Bytes
0017280
f5c65c3
154098e
6c1a39c
0017280
585c3ec
0017280
 
 
f95cdb7
585c3ec
f95cdb7
0017280
 
6c1a39c
07a69cd
9e456b7
585c3ec
9f69de7
585c3ec
c28e4d0
 
97b82b8
c28e4d0
154098e
c28e4d0
 
 
67ef019
15ee6a2
6c1a39c
0017280
 
585c3ec
 
0017280
 
 
1330e5c
 
 
 
 
0017280
 
ebd834c
 
ae2262c
 
2308eeb
 
ae2262c
 
1330e5c
 
ae2262c
 
143ed5d
 
e65bc7c
 
0017280
585c3ec
 
 
 
99b5ba2
585c3ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1479084
585c3ec
9f00d5a
46637ce
585c3ec
9bb927b
 
46637ce
9bb927b
585c3ec
46637ce
585c3ec
 
 
82a7be1
585c3ec
 
 
 
 
 
 
 
 
 
 
d23a28f
 
585c3ec
 
 
 
 
0017280
 
 
 
 
 
 
 
4ac52be
6c1a39c
4ac52be
 
 
 
6c1a39c
4ac52be
 
 
 
 
 
6c1a39c
4ac52be
 
 
6c1a39c
4ac52be
 
 
 
 
 
 
6c1a39c
4ac52be
 
 
 
 
 
 
0017280
4ac52be
 
6c1a39c
4ac52be
 
6c1a39c
4ac52be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6c1a39c
4ac52be
 
 
6c1a39c
4ac52be
 
 
 
6c1a39c
4ac52be
6c1a39c
4ac52be
 
 
 
 
 
 
 
 
6c1a39c
4ac52be
 
 
 
 
 
 
 
 
 
 
 
 
15ee6a2
 
0017280
 
 
 
 
 
 
 
 
 
 
 
2a9cbda
d0f3e65
a229e23
 
d0f3e65
0017280
e469e35
2a9cbda
 
0017280
 
 
 
 
6c1a39c
 
 
 
 
 
 
 
 
d89ddb6
 
 
 
 
4cf36f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0fa9863
4cf36f7
6c1a39c
 
 
 
 
 
3f5dbb6
720a2f4
3f5dbb6
 
 
 
 
 
 
 
b9f26b8
3f5dbb6
4cf36f7
e834d20
b54a574
 
543bce6
 
3650af5
337e106
 
6e6965b
963a61d
337e106
f1b3639
09fcd42
963a61d
337e106
963a61d
 
 
f1b3639
09fcd42
963a61d
337e106
f1b3639
09fcd42
963a61d
337e106
f1b3639
60f096b
e860365
 
 
 
 
 
 
f1b3639
337e106
f1b3639
 
e834d20
 
 
 
 
 
 
 
3650af5
3f5dbb6
2ba35cd
e068675
 
 
 
 
6c1a39c
 
 
8c2286a
6c1a39c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bfcb47b
14677f1
4cf36f7
 
61cd493
522fde6
6c1a39c
 
 
 
 
fa4f065
6c1a39c
 
 
 
 
 
 
 
fa4f065
6c1a39c
 
 
 
 
 
 
 
4cf36f7
143ed5d
 
e65bc7c
a871b98
4cf36f7
a871b98
efc5142
61cd493
6c1a39c
d89ddb6
 
0fd3783
 
e068675
ce988cc
0fd3783
 
e068675
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0fd3783
ce1923e
ce988cc
 
693929a
 
88061d3
62b61b0
0fd3783
4cf36f7
ebd834c
 
 
 
c014da5
 
 
 
 
 
 
ebd834c
 
 
 
 
2308eeb
ebd834c
 
9815e49
39c2d11
 
83087ed
ebd834c
 
3e095dc
fc60784
 
 
ebd834c
 
fc60784
ebd834c
 
 
 
9093d8f
ebd834c
2308eeb
c014da5
ebd834c
c014da5
2308eeb
9815e49
ebd834c
 
 
 
 
 
 
2308eeb
 
 
 
ebd834c
9815e49
20d594f
 
ebd834c
0017280
 
3e095dc
0017280
ebd834c
9815e49
ebd834c
 
0017280
9815e49
20d594f
9815e49
 
 
 
0017280
 
 
 
20d594f
6b689a5
20d594f
 
6b689a5
 
20d594f
 
0017280
9815e49
 
 
c23d859
 
 
 
9815e49
 
 
 
0017280
 
 
 
 
fbd18c8
0017280
9815e49
fbd18c8
0017280
 
 
 
 
 
 
 
fbd18c8
0017280
fbd18c8
0017280
 
 
e3907f4
a131ce6
0017280
a131ce6
0017280
 
 
a131ce6
0017280
 
214f864
 
 
f95cdb7
 
 
 
 
585c3ec
 
c6f476a
10eea39
585c3ec
10eea39
c6f476a
 
1479084
c6f476a
 
 
 
 
21d9988
c6f476a
21d9988
c6f476a
 
 
 
 
 
0b71bf8
c6f476a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1479084
c6f476a
 
 
 
 
 
 
 
9e456b7
 
 
 
 
 
 
 
 
 
 
 
 
 
585c3ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9e456b7
 
 
 
10eea39
9e456b7
d23a28f
9e456b7
 
 
 
 
 
 
 
 
 
585c3ec
 
 
 
10eea39
585c3ec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0017280
 
 
 
 
1330e5c
0017280
 
 
 
1330e5c
 
 
 
 
0017280
 
7e9f52f
89e2c52
7e9f52f
f95cdb7
7e9f52f
 
89e2c52
7e9f52f
c7e7022
 
 
6c1a39c
585c3ec
 
 
 
 
 
 
0017280
 
 
 
 
 
 
 
 
 
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
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
"""Streamlit web UI for the RAG Research Chatbot."""
import os
import shutil
import zipfile
import streamlit as st
import pandas as pd

from src.config_loader import load_config, get_api_key
from src.ingest import get_chroma_collection, ingest_documents
from src.kb_meta import load_kb_meta_brief
from src.query_engine import understand_query, categorize_query, init_query, _parse_vc_analyzer_result
from src.retriever import clear_collection_cache, retrieve
from src.verifier import verify_and_respond
from src.llm import list_models
from src.stata import explain_
from src.peri import APPROVED_OPTIONS
from src.peri.value_chains import categorize_vc_production, analyze_committment_to_vc , analyze_vc_committment
from src.peri.investments import analyze_investments
from src.prompts import dict_to_string
# from ui import render_db
from huggingface_hub import snapshot_download

snapshot_download(repo_id="CGIAR/peri-kb", 
                  repo_type="dataset", 
                  allow_patterns="ug/*",
                  token=os.getenv('HF_TOKEN'), 
                  local_dir="./"
                  )
shutil.copytree("./ug", "./", dirs_exist_ok=True)
# ── Session initialisation ───────────────────────────────────────────────────

def init_session():
    """Initialise st.session_state with messages list, cfg, and last_retrieval."""
    if "bot_launched" not in st.session_state:
        st.session_state.bot_launched = False
    if "messages" not in st.session_state:
        st.session_state.messages = []
    if "cfg" not in st.session_state:
        try:
            st.session_state.cfg = load_config()
        except Exception as e:
            st.error(f"Configuration error: {e}\n\nPlease run `python setup.py` first.")
            st.stop()
    if "last_retrieval" not in st.session_state:
        st.session_state.last_retrieval = None
    if "pending_clarification" not in st.session_state:
        st.session_state.pending_clarification = None
    if "unresolved_category" not in st.session_state:
        st.session_state.unresolved_category = None
    if "clarification_rounds" not in st.session_state:
        st.session_state.clarification_rounds = 0
    if "resolution_rounds" not in st.session_state:
        st.session_state.resolution_rounds = 0
    if "pending_clarification_question" not in st.session_state:
        st.session_state.pending_clarification_question = None
    if "unresolved_category_question" not in st.session_state:
        st.session_state.unresolved_category_question = None
    if "stata_options" not in st.session_state:
        st.session_state.stata_options = None
    if "stata_file" not in st.session_state:
        st.session_state.stata_file = None

# ── Sidebar ──────────────────────────────────────────────────────────────────
def render_landing_page():
    cfg = st.session_state.cfg
    # print(cfg)
    st.title(f"πŸ€– Welcome to the {cfg['chatbot'].get('name')} Assistant!")
    st.markdown("Let's go through some preliminary steps below before getting started.")
    
    st.divider()
    
    # Dropdown Options
    country = st.selectbox(
        "Please select a country you would like to prepare a PERI analysis for today:",
        [None]+[c.capitalize() for c in APPROVED_OPTIONS.get("countries")],
        format_func=lambda x: "Please select a country..." if x is None else x
    )

    regime = None
    if country is not None:
        regime = st.selectbox(
            f"Please specify whether {country.capitalize()} is an autocracy or democracy:",
            [None, "Autocracy", "Democracy"],
            format_func=lambda x: f"Please specify {country.capitalize()}'s regime..." if x is None else x
        )
    
    # creativity_level = st.select_slider(
    #     "Adjust Creativity (Temperature):",
    #     options=["Low (Factual)", "Medium (Balanced)", "High (Creative)"],
    #     value="Medium (Balanced)"
    # )
    
    uploaded_ref = None
    if regime is not None:
        # Create the file uploader widget
        pol_geo_desc = """\
            A PERI analysis typically requires a preliminary analysis of 
            the political geography of the country under review. Please upload 
            the political geography analysis results for {country} if available.
            You can also follow the guide at ... to prepare the analysis and upload your results:
            """
        st.write(pol_geo_desc.format(country=country))
        uploaded_ref = st.file_uploader(
            f"Please upload the political geography analysis results for {country}", 
            type=["csv", "xls", "xlsx"])

    if uploaded_ref is not None:
        st.divider()
        
        # Launch Button
        if st.button("Launch PERI-AI πŸš€", type="primary", use_container_width=True):
            # Save selections to session state to use during the chat session
            try:
                pol_geo_ref = pd.read_csv(uploaded_ref)
            except Exception:
                pol_geo_ref = pd.read_excel(uploaded_ref)
                
            st.session_state.country = country
            st.session_state.regime = regime
            st.session_state.pol_geo_ref = pol_geo_ref

            with st.status("Preparing metadata for the analysis...", expanded=True) as status:
                status.update(label=f'Preparing value chain metadata...')
                pol_geo_scores = categorize_vc_production(pol_geo_ref, "value_chain")
                
                status.update(label=f'Preparing investment area metadata...')
                investments = analyze_investments()

                c2vc = {}
                for vc in pol_geo_ref["value_chain"].unique():
                    status.update(label=f'Analyzing committment to {vc} production...')
                    prompt = f"Look up and retrieve all relevant information on {vc} production in {country} from the knowledgebase."     
                    qu_result = understand_query(prompt, cfg)

                    search_query = qu_result.get("search_query", prompt)
                    display_query = qu_result.get("display_query", prompt)
                    sql_query = qu_result.get("sql_query")
                    print(sql_query)

                    retrieval_result = retrieve(search_query, cfg, route="both", sql_query=sql_query)
                    st.session_state.last_retrieval = retrieval_result

                    output = verify_and_respond(
                        display_query, retrieval_result, cfg, vc, prompt,
                    )
                    print(output["response"])
                    c2vc.update(_parse_vc_analyzer_result(output["response"]))

                status.update(
                    label=f"Completed metadata preparation for {pol_geo_scores.shape[0]} value chains and {investments.shape[0]} investment areas in {country}.",
                    state="complete",
                    )
            print(c2vc)
            st.session_state.beneficiaries_var = pol_geo_scores
            st.session_state.investments_var = investments
            st.session_state.c2vc = c2vc
            st.session_state.c2vc_results = analyze_vc_committment(c2vc)
            st.session_state.country = country
            st.session_state.bot_launched = True
            st.rerun()



# ── Sidebar ──────────────────────────────────────────────────────────────────

def render_sidebar():
    """Render sidebar with provider/model selectors, web search toggle, and KB stats."""
    cfg = st.session_state.cfg

    with st.sidebar:
        st.header("Settings")

        # --- Provider dropdown ---
        providers = ["openai", "anthropic", "gemini", "meta-llama"]
        current_provider = cfg.get("llm", {}).get("provider", "openai")
        provider_index = providers.index(current_provider) if current_provider in providers else 0

        provider = st.selectbox(
            "LLM Provider",
            providers,
            index=provider_index,
            key="sidebar_provider",
        )

        # Update cfg in session when provider changes
        if provider != cfg.get("llm", {}).get("provider"):
            cfg.setdefault("llm", {})["provider"] = provider

        # --- Model dropdown (cached per provider) ---
        # Invalidate model cache if provider changed
        prev_provider_key = "prev_provider"
        if st.session_state.get(prev_provider_key) != provider:
            for p in providers:
                st.session_state.pop(f"models_{p}", None)
            st.session_state[prev_provider_key] = provider

        models_cache_key = f"models_{provider}"
        if models_cache_key not in st.session_state:
            api_key = get_api_key(cfg, provider)
            if api_key:
                try:
                    st.session_state[models_cache_key] = list_models(provider, api_key)
                except Exception:
                    st.session_state[models_cache_key] = []
            else:
                st.session_state[models_cache_key] = []

        available_models = st.session_state[models_cache_key]
        current_model = cfg.get("llm", {}).get("model", "")

        if available_models:
            model_index = (
                available_models.index(current_model)
                if current_model in available_models
                else 0
            )
            model = st.selectbox(
                "Model",
                available_models,
                index=model_index,
                key="sidebar_model",
            )
        else:
            model = st.text_input(
                "Model",
                value=current_model,
                key="sidebar_model_text",
            )

        # Update cfg in session when model changes
        if model != cfg.get("llm", {}).get("model"):
            cfg.setdefault("llm", {})["model"] = model

        # --- Web search toggle ---
        web_enabled = cfg.get("web_search", {}).get("enabled", False)
        web_toggle = st.toggle("Web search", value=web_enabled, key="sidebar_web_search")
        cfg.setdefault("web_search", {})["enabled"] = web_toggle

        st.divider()

        # --- Knowledge base stats ---
        st.subheader("Knowledge Base")
        try:
            collection = get_chroma_collection(cfg)
            chunk_count = collection.count()
            st.metric("Chunks indexed", chunk_count)
        except Exception as e:
            st.warning(f"Could not read knowledge base: {e}")
            chunk_count = 0

        # --- Re-ingest button ---
        if st.button("Re-ingest documents", use_container_width=True):
            st.info("Ingestion may take a few minutes for large document collections...")
            with st.spinner("Ingesting documents..."):
                try:
                    count = ingest_documents(cfg)
                    clear_collection_cache()
                    st.success(f"Ingested {count} chunks.")
                    # Clear cached data so it refreshes after re-ingest
                    st.session_state.pop("kb_welcome_summary", None)
                    st.rerun()
                except Exception as e:
                    st.error(f"Ingestion failed: {e}")


# ── Chat interface ───────────────────────────────────────────────────────────

def render_chat():
    """Render the chat interface with message history and input."""
    cfg = st.session_state.cfg

    # Display chat history
    for msg in st.session_state.messages:
        with st.chat_message(msg["role"]):
            st.markdown(msg["content"])

    # Chat input
    user_input = st.chat_input("Ask a question about your knowledge base...",
                            accept_file="multiple", 
                            file_type=["docx", "csv", "xlsx", "xls", "pdf", "rds", "rda", 
                                       "tsv", "sav", "dta", "txt", "md", "json", "do"]
                        )

    if user_input:
        prompt = user_input.text
        uploaded_files = user_input.files
        # Show and store user message
        st.session_state.messages.append({"role": "user", "content": prompt})
        with st.chat_message("user"):
            st.markdown(prompt)

        # ── Categorize user query ───────────────────
        save_dir = "uploaded_do_files"
        os.makedirs(save_dir, exist_ok=True)

        # Exclude the just-appended user message to avoid sending
        # the current question twice (once in history, once as query)
        cat_cfg = cfg.get("query_categorization", {})
        max_history = cat_cfg.get("max_history", 6)

        # ── Query understanding ─────────────────────────────────────────
        qu_cfg = cfg.get("query_understanding", {})
        qu_enabled = qu_cfg.get("enabled", True)
        max_history = qu_cfg.get("max_history", 6)

        # ── Check if this is a clarification response ───────────────────
        unresolved = st.session_state.unresolved_category
        if unresolved is not None:
            # This prompt is the user's clarification answer
            # Include the clarification question for context
            unresolved_question = st.session_state.get("unresolved_category_question", "")
            if unresolved_question:
                combined_cat = f"{unresolved} (Clarification: Q: {unresolved_question} A: {prompt})"
            else:
                combined_cat = f"{unresolved} β€” {prompt}"
            st.session_state.unresolved_category_question = None
            original_query_cat = unresolved
            st.session_state.unresolved_category = None
        else:
            combined_cat = prompt
            original_query_cat = prompt
            st.session_state.resolution_rounds = 0
            unresolved_question = []

        prior_messages = st.session_state.messages[:-1]
        history = [
            {"role": m["role"], "content": m["content"]}
            for m in prior_messages[-max_history:]
        ]
        try:
            qinit_result = init_query(combined_cat, cfg, history)
            print("result: ", qinit_result)
            if qinit_result.get("country", None) is None:
                
                with st.chat_message("assistant"):
                    st.markdown("Please specify the country for which you'd like to run a PERI analysis")
                return
            if qinit_result.get("value_chain", None) is None and qinit_result.get("investment", None) is None:
                
                with st.chat_message("assistant"):
                    st.markdown(f"Please specify the value chain or investment area in {qinit_result.get('country', None)} for which you'd like to run a PERI analysis")
                return
            qcat_result = categorize_query(combined_cat, cfg, history)
            print("result: ", qcat_result)
        except Exception as e:
            print(e)
            qinit_result = {"country": None, "value_chain": None, "investment": None}
            qcat_result = {"category": "pillar_1", "action": "unresolved"}

        esc_char = "\n"
        esc_char1 = "\n -"
        if not isinstance(qinit_result.get("country", None), list) and qinit_result.get("country", None)==None:
            resolution_msg = f"It seems there is no country specified for the analysis. Currently the PERI framework supports analysis for the countries listed below:\n\n  {', '.join([c.capitalize() for c in APPROVED_OPTIONS.get('countries')])}\n\n We are also continuously \
                            working to expand the framework and you can submit a form if the country you would like to run the analysis on is not included. In the meantime please let me know if you would like to run the anlysis for one of the included countries."
                        
        if qinit_result.get("country", None)[0].lower() not in [c.lower() for c in APPROVED_OPTIONS.get('countries')]:
            resolution_msg = f"It seems you are trying to run a PERI analysis for {qinit_result.get('country', None)[0].capitalize()}! Currently the PERI framework only supports analysis for the countries listed below:\n\n  {', '.join([c.capitalize() for c in APPROVED_OPTIONS.get('countries')])}\n\n We are continuously \
                working to expand the framework and you can submit a form to request. In the meantime please let me know if you would like to run the anlysis for one of the included countries."

        if qinit_result.get("value_chain", None)[0]==None and qinit_result.get('investment', None)[0]==None:
            resolution_msg = f"It seems there is no value chain or investment area specified for the analysis. Please select one of the value chains or investment areas included in the current PERI framework."
        
        if qinit_result.get("value_chain", None)[0].lower() not in [c.lower() for v in APPROVED_OPTIONS.get('value_chains').values() for c in v]:
            resolution_msg = f"It seems you are trying to run a PERI analysis for {qinit_result.get('value_chain', None)[0]} in {qinit_result.get('country', None)[0].capitalize()}! Currently the PERI framework only supports analysis for the value chains listed below:\n\n  {dict_to_string(APPROVED_OPTIONS.get('value_chains'), 2)}\n\n We are continuously \
                working to expand the framework and you can submit a form to request. In the meantime please let me know if you would like to run the anlysis for one of the included countries."
                
        if qinit_result.get("investment", None)[0].lower() not in [i.lower() for i in APPROVED_OPTIONS.get('investments')]:
            resolution_msg = f"It seems you are trying to run a PERI analysis for {qinit_result.get('investment', None)[0]} investemnt in {qinit_result.get('country', None)[0].capitalize()}! Currently the PERI framework only supports analysis for the investment areas listed below:\n\n  {', '.join([c.capitalize() for c in APPROVED_OPTIONS.get('investments')])}\n\n We are continuously \
                working to expand the framework and you can submit a form to request. In the meantime please let me know if you would like to run the anlysis for one of the included countries."

        if qinit_result.get("country", None)[0] in APPROVED_OPTIONS.get('countries') and (qinit_result.get('value_chain', None)[0] in APPROVED_OPTIONS.get('value_chains') or qinit_result.get('investment', None)[0] in APPROVED_OPTIONS.get('investments')):
    
            pillar_dict = {
                "pillar_1":f"would like to understand whether {','.join(qinit_result.get('value_chain'))} aligns with the government's political incentives.",
                "pillar_2":f"would like to understand to what degree decisions are impacted by the lobbying of particular groups or by elite influence",
                "pillar_3":f"would like to understand if {','.join(qinit_result.get('value_chain'))} and/or investing in {','.join(qinit_result.get('investment'))} can be feasibly implemented given the broader institutional and policy environment",
                
            }
            resolution_msg = f"**Before I search, could you clarify?** Please let me know if you {' and '.join([pillar_dict[c] for c in qcat_result.get('category')])}."
                                 
        # st.session_state.unresolved_category_question = f"Please let me know if you {' and '.join([pillar_dict[c] for c in qcat_result.get('category')])}."
        st.session_state.unresolved_category_question = resolution_msg
        st.session_state.unresolved_category = original_query_cat
        st.session_state.resolution_rounds += 1
        st.session_state.messages.append({"role": "assistant", "content": resolution_msg})
        
        with st.chat_message("assistant"):
            st.markdown(resolution_msg)
        return

        print("unresolved: ", unresolved)
        print("result: ", qinit_result)
        print("combined: ", combined_cat)
        # if qcat_result.get("category") == "contact_and_info" and qcat_result.get("action") == "contact":
        #     with st.chat_message("assistant"):
        #         st.markdown("For additional assistance, please contact the WEAI Helpdesk at IFPRI-WEAI@cgiar.org.\
        #                      Please let me know if there's anything else I can assist you with today.")
        #     return
        if qcat_result.get("category") == "contact_and_info" and qcat_result.get("action") == "upload":#st.session_state.clarification_rounds < max_clarifications:

            # Allow multiple file uploads
            # uploaded_files = st.file_uploader("Choose files to zip", accept_multiple_files=True)

            if uploaded_files:
                # Specify the path where the zip file will be saved
                zip_filename = "uploaded_files.zip"
                
                # Write uploaded files into a single zip archive on disk
                with zipfile.ZipFile(zip_filename, "w", zipfile.ZIP_DEFLATED) as zipf:
                    for file in uploaded_files:
                        # Write each file's bytes directly into the zip
                        zipf.writestr(file.name, file.getvalue())
                        
                st.success(f"Successfully uploaded {len(uploaded_files)} files!")

            with st.chat_message("assistant"):
                st.markdown("Thank you for sharing these resources with us. The WEAI team will work to \
                            review the files and reach out to you should we need any more information. \
                            Please let me know if there's anything else I can assist you with today.")
            return
            
        if qcat_result.get("category") == "stata" and "unresolved" in qcat_result.get("sub-action") and st.session_state.stata_options is None:
            # Ask clarification β€” store original query and question for context
            st.session_state.unresolved_category = original_query_cat

            file_path = ""
            if qcat_result.get("action") == "do":
                options = {"rewrite do-file":"rewrite", 
                           "explain the do-file":"explain", 
                           "suggest fix(es) for any errors in the do-file":"suggestfix"}
                # Accept .do files
                uploaded_file = uploaded_files[0]
                if uploaded_file is not None:
                    # Define the full file path on disk
                    file_path = os.path.join(save_dir, uploaded_file.name)

                    # Write the uploaded bytes to the local path
                    with open(file_path, "wb") as f:
                        f.write(uploaded_file.getbuffer())

                    # st.success(f"File successfully uploaded")

            if qcat_result.get("action") == "code":
                options = {"rewrite the stata code":"rewrite", 
                           "explain the stata code":"explain", 
                           "suggest fix(es) for any errors in the code":"suggestfix"}

            if qcat_result.get("action") == "error":
                options = {"explain the error":"explain", 
                           "suggest fix(es) for the error":"suggestfix"}        

            st.session_state.stata_options = options
            st.session_state.stata_file = file_path
            st.session_state.unresolved_category_question = f"Please let me know what you need help with: {' or '.join(options.keys())}. List all options that apply."
            st.session_state.resolution_rounds += 1
            resolution_msg = f"**Before I search, could you clarify?** Please let me know what you need help with: {' or '.join(options.keys())}. List all options that apply."
            st.session_state.messages.append({"role": "assistant", "content": resolution_msg})
            
            with st.chat_message("assistant"):
                st.markdown(resolution_msg)
            return

        print(qcat_result.get("sub-action"))
        if "unresolved" not in qcat_result.get("sub-action", ["unresolved"]):
            st.session_state.unresolved_category = original_query_cat
            st.session_state.resolution_rounds += 1
            st.session_state.messages.append({"role": "assistant", "content": resolution_msg})

            if qcat_result.get("category") == "stata":
                selected_options = [st.session_state.stata_options.get(s) for s in qcat_result.get("sub-action", ["unresolved"])]
                opts = {
                        "rewrite":    True if "rewrite" in selected_options else False,
                        "explain":    True if "explain" in selected_options else False,
                        "suggestfix": True if "suggestfix" in selected_options else False,
                        # "capture":    True if "capture" in selected_options else False,
                        # "verbose":    True if "verbose" in selected_options else False,
                        "lines":      None#lines if "lines" in selected_options else None
                    }
                explanation = explain_(qcat_result.get("action"), combined_cat, cfg, st.session_state.stata_file, opts)
                st.session_state.unresolved_category_question = f"{explanation}\n\n Please let me know if this answer is helpful or if it requires further clarification."
                resolution_msg = f"{explanation}\n\n Please let me know if this answer is helpful or if it requires further clarification."

                with st.chat_message("assistant"):
                    st.markdown(resolution_msg)
                return

        if "unresolved" in qcat_result.get("sub-action", "unresolved") and "Please let me know if this answer is helpful or if it requires further clarification." in unresolved_question:
            st.session_state.pending_clarification = unresolved
            st.session_state.pending_clarification = unresolved_question
            st.session_state.unresolved_category = unresolved
            st.session_state.unresolved_category_question = unresolved_question
            # qu_enabled = False

        # st.markdown(result)
            
        # ── Check if this is a clarification response ───────────────────
        pending = st.session_state.pending_clarification
        if pending is not None:
            # This prompt is the user's clarification answer
            # Include the clarification question for context
            pending_question = st.session_state.get("pending_clarification_question", "")
            if pending_question:
                combined = f"{pending} (Clarification: Q: {pending_question} A: {prompt})"
            else:
                combined = f"{pending} β€” {prompt}"
            st.session_state.pending_clarification_question = None
            original_query = pending
            st.session_state.pending_clarification = None
        else:
            combined = prompt
            original_query = prompt
            st.session_state.clarification_rounds = 0

        search_query = combined
        display_query = combined
        route = "vector"
        sql_query = None
        max_clarifications = qu_cfg.get("max_clarifications", 1)

        if qu_enabled:
            print("Q understanding")
            # Exclude the just-appended user message to avoid sending
            # the current question twice (once in history, once as query)
            prior_messages = st.session_state.messages[:-1]
            history = [
                {"role": m["role"], "content": m["content"]}
                for m in prior_messages[-max_history:]
            ]
            try:
                qu_result = understand_query(combined, cfg, history)
            except Exception:
                qu_result = {"action": "search", "search_query": combined, "display_query": combined, "original_query": original_query, "route": "vector", "sql_query": None}

            if qu_result.get("action") == "clarify" and st.session_state.clarification_rounds < max_clarifications:
                # Ask clarification β€” store original query and question for context
                st.session_state.pending_clarification = original_query
                st.session_state.pending_clarification_question = qu_result.get('clarification_question', 'Could you be more specific?')
                st.session_state.clarification_rounds += 1
                clarification_msg = f"**Before I search, could you clarify?** {qu_result.get('clarification_question', 'Could you be more specific?')}"
                st.session_state.messages.append(
                    {"role": "assistant", "content": clarification_msg}
                )
                with st.chat_message("assistant"):
                    st.markdown(clarification_msg)
                return

            # After max clarification rounds, force search (matches CLI behavior)
            if qu_result.get("action") == "clarify":
                qu_result["action"] = "search"

            search_query = qu_result.get("search_query", combined)
            display_query = qu_result.get("display_query", original_query)
            route = qu_result.get("route", "vector")
            sql_query = qu_result.get("sql_query")

        # Generate assistant response
        with st.chat_message("assistant"):
            print("Searching knowledge base")
            with st.status("Searching knowledge base...", expanded=True) as status:
                # Show reformulated query if different
                if search_query != original_query:
                    status.update(label=f'Searching for: "{search_query}"...')

                # Retrieval
                try:
                    retrieval_result = retrieve(search_query, cfg, route=route, sql_query=sql_query)
                except Exception as e:
                    status.update(label=f"Retrieval error: {e}", state="error")
                    st.error(f"Retrieval failed: {e}")
                    return
                st.session_state.last_retrieval = retrieval_result

                n_local = len(retrieval_result.get("db_results", []))
                n_web = len(retrieval_result.get("web_results", []))
                n_sql = len(retrieval_result.get("sql_results", []))
                sql_match = retrieval_result.get("sql_match_type", "")
                source_label = f"Found {n_local} local"
                if n_sql:
                    match_label = f" ({sql_match} match)" if sql_match else ""
                    source_label += f" + {n_sql} SQL rows{match_label}"
                source_label += f" + {n_web} web sources. Generating response..."
                status.update(label=source_label)

                # Response generation uses display_query β€” a clear,
                # complete question that incorporates any clarification context
                try:
                    result = verify_and_respond(
                        display_query, retrieval_result, cfg,
                        original_query=original_query,
                    )
                except Exception as e:
                    status.update(label=f"Generation error: {e}", state="error")
                    st.error(f"Response generation failed: {e}")
                    return

                # Update status based on verification outcome
                if result.get("refused"):
                    status.update(label="No sufficient sources found.", state="error")
                elif result.get("verification_passed") is True:
                    sql_label = f" + {n_sql} SQL rows" if n_sql else ""
                    status.update(
                        label=f"Verified ({result.get('iterations', 0)} iteration(s)). "
                              f"{n_local} local{sql_label} + {n_web} web sources.",
                        state="complete",
                    )
                elif result.get("verification_passed") is False:
                    status.update(
                        label="Response generated (verification did not fully pass).",
                        state="error",
                    )
                else:
                    sql_label2 = f" + {n_sql} SQL rows" if n_sql else ""
                    status.update(
                        label=f"Done. {n_local} local{sql_label2} + {n_web} web sources.",
                        state="complete",
                    )

            final_response = f"{result.get('response', '')}\n\n For additional assistance, please contact the WEAI Helpdesk at IFPRI-WEAI@cgiar.org.\
                             Please let me know if there's anything else I can assist you with today."
            # Display the response
            st.markdown(final_response)

        # Store assistant message
        st.session_state.messages.append(
            {"role": "assistant", "content": final_response}
        )

        st.session_state.unresolved_category = None
        st.session_state.unresolved_category_question = None
        
        # Cap message history to prevent unbounded memory growth
        max_messages = 200  # 100 Q&A pairs
        if len(st.session_state.messages) > max_messages:
            st.session_state.messages = st.session_state.messages[-max_messages:]

def render_db():
    # rows = st.columns((5,5), gap='medium')
    row0 = st.columns((4,4), gap='medium')
    row1 = st.columns((2,6), gap='medium')
    row2 = st.container()
    row3 = st.container()
    
    with row0[0]:
        st.markdown(f"#### Pillar 1: Value Chain Alignment with Political Incentives")       
        score_cols = ["value_chain", "trend", "strategic_importance_score", 
              "institutional_committment_score", "trend_normalized",
              "strategic_importance_normalized", "institutional_commitment_normalized"]
        
        vc_rename_cols = {
            "value_chain": "Value Chain",
            "noremalized_scores":"Distribution of Beneficiaries",
            "committment_to_vc":"Commitment to Value Chain",
            }
        
        vc_alignment = st.session_state.beneficiaries_var.merge(st.session_state.c2vc_results[score_cols], on="value_chain", how="outer")
        vc_alignment["committment_to_vc"] = vc_alignment[["strategic_importance_normalized", "institutional_commitment_normalized"]].mean(axis=1)
        vc_alignment["Alignment"] = vc_alignment[["noremalized_scores", "committment_to_vc"]].mean(axis=1)
        
        st.dataframe(vc_alignment[["value_chain", "noremalized_scores", "committment_to_vc", "Alignment"]].rename(columns=vc_rename_cols),
                    # column_order=("states", "population"),
                    hide_index=True,
                    width='stretch',
                    # column_config={
                    #     "states": st.column_config.TextColumn(
                    #         "States",
                    #     ),
                    #     "population": st.column_config.ProgressColumn(
                    #         "Population",
                    #         format="%f",
                    #         min_value=0,
                    #         max_value=max(st.session_state.beneficiaries_var.population),
                    #     )}
                    )
    with row0[1]:
        st.markdown(f"#### Pillar 1: Investment Alignment with Political Incentives") 
        investment_rename_cols = {
                "investments":"Investment",
                "time_to_impact":"Time to Impact",
                "targetability":"Targetability",
                "visibility":"Visibility",
                "agg_score":"Investment",
            }
        st.dataframe(st.session_state.investments_var[["investments","time_to_impact", "targetability", "visibility"]].rename(columns=investment_rename_cols),
                    # column_order=("states", "population"),
                    hide_index=True,
                    width='stretch',
                    # column_config={
                    #     "states": st.column_config.TextColumn(
                    #         "States",
                    #     ),
                    #     "population": st.column_config.ProgressColumn(
                    #         "Population",
                    #         format="%f",
                    #         min_value=0,
                    #         max_value=max(st.session_state.beneficiaries_var.population),
                    #     )}
                    )
    with row1[0]:
        st.markdown(f'#### Political Geography of {st.session_state.country}')
        st.dataframe(st.session_state.beneficiaries_var,
                    # column_order=("states", "population"),
                    hide_index=True,
                    width='stretch',
                    # column_config={
                    #     "states": st.column_config.TextColumn(
                    #         "States",
                    #     ),
                    #     "population": st.column_config.ProgressColumn(
                    #         "Population",
                    #         format="%f",
                    #         min_value=0,
                    #         max_value=max(st.session_state.beneficiaries_var.population),
                    #     )}
                    )
            
    with row1[1]:
        st.markdown(f'#### Investment Area Analysis Results')
        st.dataframe(st.session_state.investments_var,
                    # column_order=("states", "population"),
                    hide_index=True,
                    width='stretch',
                    # column_config={
                    #     "states": st.column_config.TextColumn(
                    #         "States",
                    #     ),
                    #     "population": st.column_config.ProgressColumn(
                    #         "Population",
                    #         format="%f",
                    #         min_value=0,
                    #         max_value=max(st.session_state.beneficiaries_var.population),
                    #     )}
                    )  
    with row2:
        st.markdown(f'#### Committment to Value Chain Analysis')
        st.dataframe(st.session_state.c2vc_results,
                    # column_order=("states", "population"),
                    hide_index=True,
                    width='stretch',
                    # column_config={
                    #     "states": st.column_config.TextColumn(
                    #         "States",
                    #     ),
                    #     "population": st.column_config.ProgressColumn(
                    #         "Population",
                    #         format="%f",
                    #         min_value=0,
                    #         max_value=max(st.session_state.beneficiaries_var.population),
                    #     )}
                    )
    with row3:
        # with cols[2]:
        st.markdown(f"#### Committment of {st.session_state.country}'s Government to Various Value Chains")
       
       
        def display_nested_dict(d, indent=0, cols_per_row=2):
        # Get all items in the dictionary
            items = list(d.items())

            # Create rows of columns
            for i in range(0, len(items), cols_per_row):
                cols = st.columns(cols_per_row)
                for j in range(cols_per_row):
                    if i + j < len(items):
                        key, value = items[i + j]
                        with cols[j]:
                            if isinstance(value, dict):
                                with st.expander(f"{' ' * indent}πŸ“ {key.capitalize()}"):
                                # Recursive call (you can adjust columns or keep single column inside expander)
                                    display_nested_dict(value, indent + 2, 1)
                            else:
                                st.write(f"**{key.capitalize()}:** {value}")

        display_nested_dict(st.session_state.c2vc, cols_per_row=3)
        # st.json(st.session_state.c2vc)
# ── Main ─────────────────────────────────────────────────────────────────────

def main():
    """Orchestrate the Streamlit app."""
    st.set_page_config(
        page_title="ResearchBot",
        page_icon="πŸ”¬",
        layout="wide",
    )

    init_session()

    cfg = st.session_state.cfg
    bot_name = cfg.get("chatbot", {}).get("name", "ResearchBot")

    st.title(bot_name)

    # Show KB summary on first visit (LLM-generated welcome summary)
    if not st.session_state.messages:
        if "kb_welcome_summary" not in st.session_state:
            st.session_state.kb_welcome_summary = load_kb_meta_brief(cfg)
        kb_summary = st.session_state.kb_welcome_summary
        if kb_summary:
            with st.expander("Knowledge Base Contents", expanded=True):
                st.markdown(kb_summary)
                st.markdown("In case you are familiar with the PERI analysis \
                please let me know what pillar(s) you'd like us to analyze today and for what country. If this is\
                is your first time working with PERI please proceed by letting me know what your particular interests are.")

    if not st.session_state.bot_launched:
        render_landing_page()

    else:
        # render_sidebar()
        # render_chat()
        render_db()

    st.caption(
        "All answers are sourced from the local knowledge base. "
        "Web sources are supplementary only. "
        "Every claim is citation-verified before display."
    )


if __name__ == "__main__":
    main()