Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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);
}
}
Original file line number Diff line number Diff line change
@@ -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);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand All @@ -41,8 +41,8 @@ public class ChatbotController {
* @return list of all chats
*/
@GetMapping("/chats/{run}/{workgroup}")
public ResponseEntity<List<Chat>> getAllChats(@PathVariable RunImpl run,
@PathVariable WorkgroupImpl workgroup) {
public ResponseEntity<List<Chat>> getAllChats(@PathVariable Run run,
@PathVariable Workgroup workgroup) {
return ResponseEntity.ok(chatbotService.getAllChats(run, workgroup));
}

Expand All @@ -55,10 +55,10 @@ public ResponseEntity<List<Chat>> getAllChats(@PathVariable RunImpl run,
* @return the created chat
*/
@PostMapping("/chats/{run}/{workgroup}")
public ResponseEntity<Chat> createChat(@PathVariable RunImpl run,
@PathVariable WorkgroupImpl workgroup, @RequestBody Chat chat) {
public ResponseEntity<Chat> 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));
}

/**
Expand All @@ -72,9 +72,9 @@ public ResponseEntity<Chat> createChat(@PathVariable RunImpl run,
* @throws ObjectNotFoundException when the chat is not found
*/
@PutMapping("/chats/{run}/{workgroup}/{chatId}")
public ResponseEntity<Chat> updateChat(@PathVariable RunImpl run,
@PathVariable WorkgroupImpl workgroup, @PathVariable Long chatId, @RequestBody Chat chat)
throws ObjectNotFoundException {
public ResponseEntity<Chat> updateChat(@PathVariable Run run,
@PathVariable Workgroup workgroup, @PathVariable Long chatId,
@RequestBody Chat chat) throws ObjectNotFoundException {
return ResponseEntity.ok(chatbotService.updateChat(run, workgroup, chatId, chat));
}

Expand All @@ -88,8 +88,9 @@ public ResponseEntity<Chat> updateChat(@PathVariable RunImpl run,
* @throws ObjectNotFoundException when the chat is not found
*/
@DeleteMapping("/chats/{run}/{workgroup}/{chatId}")
public ResponseEntity<Void> deleteChat(@PathVariable RunImpl run,
@PathVariable WorkgroupImpl workgroup, @PathVariable Long chatId) throws ObjectNotFoundException {
public ResponseEntity<Void> deleteChat(@PathVariable Run run,
@PathVariable Workgroup workgroup, @PathVariable Long chatId)
throws ObjectNotFoundException {
chatbotService.deleteChat(run, workgroup, chatId);
return ResponseEntity.noContent().build();
}
Expand Down
Original file line number Diff line number Diff line change
@@ -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());
}
}
Loading
Loading