| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193 |
- 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<Map<String, Object>> 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<Map<String, Object>> docs = retrieverService.search(query, intent, 20);
- docs = rerank(docs, query, 5);
- List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
- String answer = llmService.chat(messages);
- List<Map<String, Object>> 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<ServerSentEvent<String>> 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<Map<String, Object>> docs = rerank(retrieverService.search(query, intent, 20), query, 5);
- Sinks.Many<ServerSentEvent<String>> sink = Sinks.many().unicast().onBackpressureBuffer();
- try {
- sink.tryEmitNext(ServerSentEvent.<String>builder().event("intent").data(intent).build());
- sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Retrieving...").build());
- sink.tryEmitNext(ServerSentEvent.<String>builder().event("status").data("Matched " + docs.size() + " records, generating...").build());
- final List<Map<String, String>> messages = promptService.buildPrompt(query, docs, intent);
- StringBuilder fullAnswer = new StringBuilder();
- llmService.chatStream(messages)
- .doOnNext(token -> {
- fullAnswer.append(token);
- sink.tryEmitNext(ServerSentEvent.<String>builder().data(token).build());
- })
- .doOnComplete(() -> {
- final List<Map<String, Object>> sources = buildSources(docs);
- try {
- String meta = new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(Map.of(
- "intent", intent,
- "sources", sources
- ));
- sink.tryEmitNext(ServerSentEvent.<String>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<Map<String, Object>> 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<Map<String, Object>> getConversationDetail(@PathVariable String cid) {
- var msgs = persistenceService.getConversationDetail(cid);
- return ResponseEntity.ok(Map.of("conversation_id", cid, "messages", msgs));
- }
- @PostMapping("/feedback")
- public ResponseEntity<Map<String, Object>> submitFeedback(@RequestBody FeedbackRequest request) {
- persistenceService.updateFeedback(request.getMessageId(), request.getFeedback());
- return ResponseEntity.ok(Map.of("status", "ok"));
- }
- @GetMapping("/admin/conversations")
- public ResponseEntity<Map<String, Object>> 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<Object> 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<Map<String, Object>> items = jdbc.queryForList(
- sql.toString(), params.toArray());
- return ResponseEntity.ok(Map.of(
- "items", items,
- "page", page,
- "page_size", pageSize
- ));
- }
- private List<Map<String, Object>> buildSources(List<Map<String, Object>> docs) {
- return docs.stream()
- .map(d -> Map.<String, Object>of(
- "name", d.getOrDefault("source", ""),
- "section", d.getOrDefault("section", ""),
- "source", d.getOrDefault("source", "")
- ))
- .collect(Collectors.toList());
- }
- private List<Map<String, Object>> rerank(List<Map<String, Object>> 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;
- }
- }
|