diff --git a/src/main/java/org/wise/portal/presentation/web/AWSBedrockController.java b/src/main/java/org/wise/portal/presentation/web/AWSBedrockController.java index 253c64c67..423e7c62e 100644 --- a/src/main/java/org/wise/portal/presentation/web/AWSBedrockController.java +++ b/src/main/java/org/wise/portal/presentation/web/AWSBedrockController.java @@ -1,67 +1,49 @@ package org.wise.portal.presentation.web; -import java.io.BufferedReader; -import java.io.IOException; -import java.io.InputStreamReader; -import java.io.OutputStreamWriter; -import java.net.HttpURLConnection; -import java.net.URL; - -import org.springframework.beans.factory.annotation.Autowired; import org.springframework.core.env.Environment; +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; import org.springframework.security.access.annotation.Secured; import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.ResponseBody; import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.client.RestClient; @RestController @RequestMapping("/api/aws-bedrock/chat") public class AWSBedrockController { - @Autowired - Environment appProperties; + private final Environment appProperties; + private final RestClient restClient; - @ResponseBody - @Secured("ROLE_USER") - @PostMapping - protected String sendChatMessage(@RequestBody String body) { - String apiKey = appProperties.getProperty("aws.bedrock.api.key"); - if (apiKey == null || apiKey.isEmpty()) { - throw new RuntimeException("aws.bedrock.api.key is not set"); - } - String apiEndpoint = appProperties.getProperty("aws.bedrock.runtime.endpoint"); - if (apiEndpoint == null || apiEndpoint.isEmpty()) { - throw new RuntimeException("aws.bedrock.runtime.endpoint is not set"); - } - // assume openai-only support for now. We'll add other models later. - apiEndpoint += "/openai/v1/chat/completions"; + public AWSBedrockController(Environment appProperties, RestClient.Builder restClientBuilder) { + this.appProperties = appProperties; + this.restClient = restClientBuilder.build(); + } - try { - URL url = new URL(apiEndpoint); - HttpURLConnection connection = (HttpURLConnection) url.openConnection(); - connection.setRequestMethod("POST"); - connection.setRequestProperty("Authorization", "Bearer " + apiKey); - connection.setRequestProperty("Content-Type", "application/json; charset=utf-8"); - connection.setRequestProperty("Accept-Charset", "UTF-8"); - connection.setDoOutput(true); - OutputStreamWriter writer = new OutputStreamWriter(connection.getOutputStream()); - writer.write(body); - writer.flush(); - writer.close(); - BufferedReader br = new BufferedReader( - new InputStreamReader(connection.getInputStream(), "UTF-8")); - String line; - StringBuffer response = new StringBuffer(); - while ((line = br.readLine()) != null) { - response.append(line); - } - br.close(); - return response.toString(); - } catch (IOException e) { - throw new RuntimeException(e); - } - } + @ResponseBody + @Secured("ROLE_USER") + @PostMapping(produces = "application/json;charset=UTF-8") + public String sendChatMessage(@RequestBody String body) { + String apiKey = appProperties.getProperty("aws.bedrock.api.key"); + if (apiKey == null || apiKey.isEmpty()) { + throw new RuntimeException("aws.bedrock.api.key is not set"); + } + String apiEndpoint = appProperties.getProperty("aws.bedrock.runtime.endpoint"); + if (apiEndpoint == null || apiEndpoint.isEmpty()) { + throw new RuntimeException("aws.bedrock.runtime.endpoint is not set"); + } + // assume openai-only support for now. We'll add other models later. + apiEndpoint += "/openai/v1/chat/completions"; + return restClient.post() + .uri(apiEndpoint) + .header(HttpHeaders.AUTHORIZATION, "Bearer " + apiKey) + .contentType(MediaType.APPLICATION_JSON) + .body(body) + .retrieve() + .body(String.class); + } } diff --git a/src/main/java/org/wise/portal/presentation/web/controllers/ChatGptController.java b/src/main/java/org/wise/portal/presentation/web/controllers/ChatGptController.java index ca856e63a..94dc47ca2 100644 --- a/src/main/java/org/wise/portal/presentation/web/controllers/ChatGptController.java +++ b/src/main/java/org/wise/portal/presentation/web/controllers/ChatGptController.java @@ -1,59 +1,46 @@ package org.wise.portal.presentation.web.controllers; -import java.io.BufferedReader; -import java.io.IOException; -import java.io.InputStreamReader; -import java.io.OutputStreamWriter; -import java.net.HttpURLConnection; -import java.net.URL; import org.springframework.beans.factory.annotation.Value; +import org.springframework.http.HttpHeaders; +import org.springframework.http.MediaType; import org.springframework.security.access.annotation.Secured; import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.ResponseBody; import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.client.RestClient; @RestController @RequestMapping("/api/chat-gpt") public class ChatGptController { - @Value("${openai.api.key:}") - private String openAiApiKey; + private final String openAiApiKey; + private final String openAiChatApiUrl; + private final RestClient restClient; - @Value("${openai.chat.api.url:https://api.openai.com/v1/chat/completions}") - private String openAiChatApiUrl; + public ChatGptController( + @Value("${openai.api.key:}") String openAiApiKey, + @Value("${openai.chat.api.url:https://api.openai.com/v1/chat/completions}") String openAiChatApiUrl, + RestClient.Builder restClientBuilder) { + this.openAiApiKey = openAiApiKey; + this.openAiChatApiUrl = openAiChatApiUrl; + this.restClient = restClientBuilder.build(); + } @ResponseBody @Secured("ROLE_USER") @PostMapping(produces = "application/json;charset=UTF-8") - protected String sendChatMessage(@RequestBody String body) { + public String sendChatMessage(@RequestBody String body) { if (openAiApiKey == null || openAiApiKey.isEmpty()) { throw new RuntimeException("openai.api.key is not set"); } - try { - URL url = new URL(openAiChatApiUrl); - HttpURLConnection connection = (HttpURLConnection) url.openConnection(); - connection.setRequestMethod("POST"); - connection.setRequestProperty("Authorization", "Bearer " + openAiApiKey); - connection.setRequestProperty("Content-Type", "application/json; charset=utf-8"); - connection.setRequestProperty("Accept-Charset", "UTF-8"); - connection.setDoOutput(true); - OutputStreamWriter writer = new OutputStreamWriter(connection.getOutputStream()); - writer.write(body); - writer.flush(); - writer.close(); - BufferedReader br = new BufferedReader( - new InputStreamReader(connection.getInputStream(), "UTF-8")); - String line; - StringBuffer response = new StringBuffer(); - while ((line = br.readLine()) != null) { - response.append(line); - } - br.close(); - return response.toString(); - } catch (IOException e) { - throw new RuntimeException(e); - } + return restClient.post() + .uri(openAiChatApiUrl) + .header(HttpHeaders.AUTHORIZATION, "Bearer " + openAiApiKey) + .contentType(MediaType.APPLICATION_JSON) + .body(body) + .retrieve() + .body(String.class); } } diff --git a/src/main/java/org/wise/portal/presentation/web/controllers/ChatbotController.java b/src/main/java/org/wise/portal/presentation/web/controllers/ChatbotController.java index 97cdc0bba..9f3da2f72 100644 --- a/src/main/java/org/wise/portal/presentation/web/controllers/ChatbotController.java +++ b/src/main/java/org/wise/portal/presentation/web/controllers/ChatbotController.java @@ -15,8 +15,8 @@ import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RestController; import org.wise.portal.dao.ObjectNotFoundException; -import org.wise.portal.domain.run.impl.RunImpl; -import org.wise.portal.domain.workgroup.impl.WorkgroupImpl; +import org.wise.portal.domain.run.Run; +import org.wise.portal.domain.workgroup.Workgroup; import org.wise.portal.service.chatbot.ChatbotService; import org.wise.vle.domain.chatbot.Chat; @@ -41,8 +41,8 @@ public class ChatbotController { * @return list of all chats */ @GetMapping("/chats/{run}/{workgroup}") - public ResponseEntity> getAllChats(@PathVariable RunImpl run, - @PathVariable WorkgroupImpl workgroup) { + public ResponseEntity> getAllChats(@PathVariable Run run, + @PathVariable Workgroup workgroup) { return ResponseEntity.ok(chatbotService.getAllChats(run, workgroup)); } @@ -55,10 +55,10 @@ public ResponseEntity> getAllChats(@PathVariable RunImpl run, * @return the created chat */ @PostMapping("/chats/{run}/{workgroup}") - public ResponseEntity createChat(@PathVariable RunImpl run, - @PathVariable WorkgroupImpl workgroup, @RequestBody Chat chat) { + public ResponseEntity createChat(@PathVariable Run run, + @PathVariable Workgroup workgroup, @RequestBody Chat chat) { return ResponseEntity.status(HttpStatus.CREATED) - .body(chatbotService.createChat(run, workgroup, chat)); + .body(chatbotService.createChat(run, workgroup, chat)); } /** @@ -72,9 +72,9 @@ public ResponseEntity createChat(@PathVariable RunImpl run, * @throws ObjectNotFoundException when the chat is not found */ @PutMapping("/chats/{run}/{workgroup}/{chatId}") - public ResponseEntity updateChat(@PathVariable RunImpl run, - @PathVariable WorkgroupImpl workgroup, @PathVariable Long chatId, @RequestBody Chat chat) - throws ObjectNotFoundException { + public ResponseEntity updateChat(@PathVariable Run run, + @PathVariable Workgroup workgroup, @PathVariable Long chatId, + @RequestBody Chat chat) throws ObjectNotFoundException { return ResponseEntity.ok(chatbotService.updateChat(run, workgroup, chatId, chat)); } @@ -88,8 +88,9 @@ public ResponseEntity updateChat(@PathVariable RunImpl run, * @throws ObjectNotFoundException when the chat is not found */ @DeleteMapping("/chats/{run}/{workgroup}/{chatId}") - public ResponseEntity deleteChat(@PathVariable RunImpl run, - @PathVariable WorkgroupImpl workgroup, @PathVariable Long chatId) throws ObjectNotFoundException { + public ResponseEntity deleteChat(@PathVariable Run run, + @PathVariable Workgroup workgroup, @PathVariable Long chatId) + throws ObjectNotFoundException { chatbotService.deleteChat(run, workgroup, chatId); return ResponseEntity.noContent().build(); } diff --git a/src/test/java/org/wise/portal/presentation/web/AWSBedrockControllerTest.java b/src/test/java/org/wise/portal/presentation/web/AWSBedrockControllerTest.java new file mode 100644 index 000000000..e9e41ea19 --- /dev/null +++ b/src/test/java/org/wise/portal/presentation/web/AWSBedrockControllerTest.java @@ -0,0 +1,86 @@ +package org.wise.portal.presentation.web; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.content; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.header; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.method; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo; +import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.MediaType; +import org.springframework.mock.env.MockEnvironment; +import org.springframework.test.web.client.MockRestServiceServer; +import org.springframework.web.client.RestClient; + +public class AWSBedrockControllerTest { + + private static final String API_KEY = "bedrock-key-12345"; + private static final String BASE_ENDPOINT = "https://bedrock.example.com"; + private static final String EXPECTED_URL = BASE_ENDPOINT + "/openai/v1/chat/completions"; + + private MockEnvironment env; + private MockRestServiceServer mockServer; + private AWSBedrockController controller; + + @BeforeEach + public void setUp() { + env = new MockEnvironment(); + env.setProperty("aws.bedrock.api.key", API_KEY); + env.setProperty("aws.bedrock.runtime.endpoint", BASE_ENDPOINT); + + RestClient.Builder builder = RestClient.builder(); + mockServer = MockRestServiceServer.bindTo(builder).build(); + controller = new AWSBedrockController(env, builder); + } + + @Test + public void sendChatMessage_successfulResponse() { + String requestBody = "{\"prompt\":\"test prompt\"}"; + String expectedResponse = "{\"response\":\"test response\"}"; + + mockServer.expect(requestTo(EXPECTED_URL)) + .andExpect(method(HttpMethod.POST)) + .andExpect(header(HttpHeaders.AUTHORIZATION, "Bearer " + API_KEY)) + .andExpect(content().contentType(MediaType.APPLICATION_JSON)) + .andExpect(content().string(requestBody)) + .andRespond(withSuccess(expectedResponse, MediaType.APPLICATION_JSON)); + + String result = controller.sendChatMessage(requestBody); + + assertEquals(expectedResponse, result); + mockServer.verify(); + } + + @Test + public void sendChatMessage_missingApiKey_throwsException() { + MockEnvironment missingKeyEnv = new MockEnvironment(); + missingKeyEnv.setProperty("aws.bedrock.runtime.endpoint", BASE_ENDPOINT); + RestClient.Builder builder = RestClient.builder(); + AWSBedrockController ctrl = new AWSBedrockController(missingKeyEnv, builder); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> { + ctrl.sendChatMessage("{}"); + }); + + assertEquals("aws.bedrock.api.key is not set", exception.getMessage()); + } + + @Test + public void sendChatMessage_missingEndpoint_throwsException() { + MockEnvironment missingEndpointEnv = new MockEnvironment(); + missingEndpointEnv.setProperty("aws.bedrock.api.key", API_KEY); + RestClient.Builder builder = RestClient.builder(); + AWSBedrockController ctrl = new AWSBedrockController(missingEndpointEnv, builder); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> { + ctrl.sendChatMessage("{}"); + }); + + assertEquals("aws.bedrock.runtime.endpoint is not set", exception.getMessage()); + } +} diff --git a/src/test/java/org/wise/portal/presentation/web/controllers/ChatGptControllerTest.java b/src/test/java/org/wise/portal/presentation/web/controllers/ChatGptControllerTest.java new file mode 100644 index 000000000..a1a11c046 --- /dev/null +++ b/src/test/java/org/wise/portal/presentation/web/controllers/ChatGptControllerTest.java @@ -0,0 +1,75 @@ +package org.wise.portal.presentation.web.controllers; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.content; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.header; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.method; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo; +import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.MediaType; +import org.springframework.test.web.client.MockRestServiceServer; +import org.springframework.web.client.RestClient; + +public class ChatGptControllerTest { + + private static final String API_KEY = "sk-test-key-12345"; + private static final String CHAT_URL = "https://api.openai.com/v1/chat/completions"; + + private MockRestServiceServer mockServer; + private ChatGptController controller; + + @BeforeEach + public void setUp() { + RestClient.Builder builder = RestClient.builder(); + mockServer = MockRestServiceServer.bindTo(builder).build(); + controller = new ChatGptController(API_KEY, CHAT_URL, builder); + } + + @Test + public void sendChatMessage_successfulResponse() { + String requestBody = "{\"model\":\"gpt-4\",\"messages\":[{\"role\":\"user\",\"content\":\"hello\"}]}"; + String expectedResponse = "{\"choices\":[{\"message\":{\"role\":\"assistant\",\"content\":\"hi!\"}}]}"; + + mockServer.expect(requestTo(CHAT_URL)) + .andExpect(method(HttpMethod.POST)) + .andExpect(header(HttpHeaders.AUTHORIZATION, "Bearer " + API_KEY)) + .andExpect(content().contentType(MediaType.APPLICATION_JSON)) + .andExpect(content().string(requestBody)) + .andRespond(withSuccess(expectedResponse, MediaType.APPLICATION_JSON)); + + String result = controller.sendChatMessage(requestBody); + + assertEquals(expectedResponse, result); + mockServer.verify(); + } + + @Test + public void sendChatMessage_missingApiKey_throwsException() { + RestClient.Builder builder = RestClient.builder(); + ChatGptController controllerWithoutKey = new ChatGptController("", CHAT_URL, builder); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> { + controllerWithoutKey.sendChatMessage("{\"messages\":[]}"); + }); + + assertEquals("openai.api.key is not set", exception.getMessage()); + } + + @Test + public void sendChatMessage_nullApiKey_throwsException() { + RestClient.Builder builder = RestClient.builder(); + ChatGptController controllerWithoutKey = new ChatGptController(null, CHAT_URL, builder); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> { + controllerWithoutKey.sendChatMessage("{\"messages\":[]}"); + }); + + assertEquals("openai.api.key is not set", exception.getMessage()); + } +}