package com.pharmacopoeia.controller; import com.pharmacopoeia.dto.ChatRequest; import com.pharmacopoeia.dto.FeedbackRequest; import com.pharmacopoeia.service.*; import org.springframework.http.MediaType; import org.springframework.http.ResponseEntity; import org.springframework.http.codec.ServerSentEvent; import org.springframework.web.bind.annotation.*; import reactor.core.publisher.Flux; import reactor.core.publisher.Sinks; import org.springframework.jdbc.core.JdbcTemplate; import java.util.*; import java.util.stream.Collectors; @RestController @RequestMapping("/api/v1/chat") public class ChatController { private final RetrieverService retrieverService; private final LLMService llmService; private final PromptService promptService; private final ChatPersistenceService persistenceService; private final JdbcTemplate jdbc; public ChatController(RetrieverService rs, LLMService ls, PromptService ps, ChatPersistenceService cps, JdbcTemplate jdbc) { this.retrieverService = rs; this.llmService = ls; this.promptService = ps; this.persistenceService = cps; this.jdbc = jdbc; } @PostMapping("/ask") public ResponseEntity> chatAsk(@RequestBody ChatRequest request) { String query = request.getMessage(); String cid = request.getConversationId() != null && !request.getConversationId().isBlank() ? request.getConversationId() : UUID.randomUUID().toString(); String intent = retrieverService.classifyIntent(query); List> docs = retrieverService.search(query, intent, 20); docs = rerank(docs, query, 5); List> messages = promptService.buildPrompt(query, docs, intent); String answer = llmService.chat(messages); List> sources = buildSources(docs); persistenceService.saveMessage(cid, "user", query, intent, null); persistenceService.saveMessage(cid, "assistant", answer, intent, sources); return ResponseEntity.ok(Map.of( "answer", answer, "sources", sources, "intent", intent )); } @PostMapping(value = "/stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE) public Flux> chatStream(@RequestBody ChatRequest request) { String query = request.getMessage(); final String cid = request.getConversationId() != null && !request.getConversationId().isBlank() ? request.getConversationId() : UUID.randomUUID().toString(); final String intent = retrieverService.classifyIntent(query); final List> docs = rerank(retrieverService.search(query, intent, 20), query, 5); Sinks.Many> sink = Sinks.many().unicast().onBackpressureBuffer(); try { sink.tryEmitNext(ServerSentEvent.builder().event("intent").data(intent).build()); sink.tryEmitNext(ServerSentEvent.builder().event("status").data("Retrieving...").build()); sink.tryEmitNext(ServerSentEvent.builder().event("status").data("Matched " + docs.size() + " records, generating...").build()); final List> messages = promptService.buildPrompt(query, docs, intent); StringBuilder fullAnswer = new StringBuilder(); llmService.chatStream(messages) .doOnNext(token -> { fullAnswer.append(token); sink.tryEmitNext(ServerSentEvent.builder().data(token).build()); }) .doOnComplete(() -> { final List> sources = buildSources(docs); try { String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of( "intent", intent, "sources", sources )); sink.tryEmitNext(ServerSentEvent.builder().event("meta").data(meta).build()); } catch (Exception ignored) {} persistenceService.saveMessage(cid, "user", query, intent, null); persistenceService.saveMessage(cid, "assistant", fullAnswer.toString(), intent, sources); sink.tryEmitComplete(); }) .doOnError(e -> sink.tryEmitError(e)) .subscribe(); } catch (Exception e) { sink.tryEmitError(e); } return sink.asFlux(); } @GetMapping("/history") public ResponseEntity> getHistory( @RequestParam(defaultValue = "1") int page, @RequestParam(defaultValue = "20") int pageSize) { var items = persistenceService.getHistory(page, pageSize); return ResponseEntity.ok(Map.of("items", items, "page", page, "page_size", pageSize)); } @GetMapping("/history/{cid}") public ResponseEntity> getConversationDetail(@PathVariable String cid) { var msgs = persistenceService.getConversationDetail(cid); return ResponseEntity.ok(Map.of("conversation_id", cid, "messages", msgs)); } @PostMapping("/feedback") public ResponseEntity> submitFeedback(@RequestBody FeedbackRequest request) { persistenceService.updateFeedback(request.getMessageId(), request.getFeedback()); return ResponseEntity.ok(Map.of("status", "ok")); } @GetMapping("/admin/conversations") public ResponseEntity> adminListConversations( @RequestParam(defaultValue = "1") int page, @RequestParam(defaultValue = "20") int pageSize, @RequestParam(required = false) String keyword) { int offset = (page - 1) * pageSize; StringBuilder sql = new StringBuilder(""" SELECT DISTINCT ON (c.conversation_id) c.conversation_id, c.title, c.created_at, m.content AS last_msg, m.role FROM conversations c JOIN messages m ON m.conversation_id = c.conversation_id """); List params = new ArrayList<>(); if (keyword != null && !keyword.isBlank()) { sql.append("WHERE m.content ILIKE ? "); params.add("%" + keyword + "%"); } sql.append(""" ORDER BY c.conversation_id, m.created_at DESC LIMIT ? OFFSET ? """); params.add(pageSize); params.add(offset); List> items = jdbc.queryForList( sql.toString(), params.toArray()); return ResponseEntity.ok(Map.of( "items", items, "page", page, "page_size", pageSize )); } private List> buildSources(List> docs) { return docs.stream() .map(d -> Map.of( "name", d.getOrDefault("source", ""), "section", d.getOrDefault("section", ""), "source", d.getOrDefault("source", "") )) .collect(Collectors.toList()); } private List> rerank(List> docs, String query, int topK) { if (docs.size() <= topK) return docs; docs.sort((a, b) -> { double sa = toDouble(a.get("similarity")); double sb = toDouble(b.get("similarity")); return Double.compare(sb, sa); }); return docs.subList(0, Math.min(topK, docs.size())); } private double toDouble(Object o) { if (o instanceof Number n) return n.doubleValue(); return 0; } }