diff --git a/src/control-plane-services/cloud-functions/nvcf-core/src/main/java/com/nvidia/nvcf/grpc/GrpcInvocationService.java b/src/control-plane-services/cloud-functions/nvcf-core/src/main/java/com/nvidia/nvcf/grpc/GrpcInvocationService.java index 577ceadedc..ab0f171938 100644 --- a/src/control-plane-services/cloud-functions/nvcf-core/src/main/java/com/nvidia/nvcf/grpc/GrpcInvocationService.java +++ b/src/control-plane-services/cloud-functions/nvcf-core/src/main/java/com/nvidia/nvcf/grpc/GrpcInvocationService.java @@ -20,8 +20,10 @@ import static com.nvidia.nvcf.util.NvcfConstants.SCOPE_INVOKE_FUNCTION; import com.google.common.annotations.VisibleForTesting; +import com.nvidia.boot.exceptions.NotFoundException; import com.nvidia.boot.exceptions.UnauthorizedException; import com.nvidia.nvcf.configuration.exceptions.InvalidInvocationException; +import com.nvidia.nvcf.persistence.function.entity.FunctionType; import com.nvidia.nvcf.proto.ClientInvokeRequest; import com.nvidia.nvcf.proto.ClientInvokeResponse; import com.nvidia.nvcf.proto.ClientInvokeResponse.FunctionVersion; @@ -46,6 +48,10 @@ @RequiredArgsConstructor public class GrpcInvocationService extends InvocationImplBase { + private static final String MESG_LLM_FUNCTION_NOT_INVOCABLE = + "Function id '%s': LLM functions cannot be invoked through this endpoint. " + + "Use the LLM API instead."; + private final GrpcAuthService grpcAuthService; private final AccountService accountService; private final FunctionInvocationValidationService functionInvocationValidationService; @@ -64,10 +70,14 @@ public void authClientInvocation( UUID.fromString(request.getFunctionVersionId()) : null; var ncaId = request.hasTargetNcaId() ? request.getTargetNcaId() : accountService.getNcaId(authentication); - var functions = lookupAndValidateAccess(authentication, - ncaId, + // LLM workers do not consume the request queues this path publishes to, so a request + // would wait out the poll window and time out. Only return versions that can serve it. + var functions = withoutLlmVersions(ncaId, functionId, - functionVersionId); + lookupAndValidateAccess(authentication, + ncaId, + functionId, + functionVersionId)); // picking the first function version for function level info. // all function versions will be of the same function. @@ -154,4 +164,20 @@ private List lookupAndValidateAccess( } } + private static List withoutLlmVersions( + String ncaId, + UUID functionId, + List functions) { + var invocable = functions.stream() + .filter(context -> context.targetFunction().getFunctionType() != FunctionType.LLM) + .toList(); + if (invocable.isEmpty()) { + var mesg = MESG_LLM_FUNCTION_NOT_INVOCABLE.formatted(functionId); + log.warn("Rejecting classic invocation of LLM function: ncaId={}, functionId={}", + ncaId, functionId); + throw new InvalidInvocationException(ncaId, new NotFoundException(mesg)); + } + return invocable; + } + } diff --git a/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/grpc/GrpcInvocationServiceTest.java b/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/grpc/GrpcInvocationServiceTest.java index 6015dfd72c..f44f00b4d6 100644 --- a/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/grpc/GrpcInvocationServiceTest.java +++ b/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/grpc/GrpcInvocationServiceTest.java @@ -44,9 +44,11 @@ import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_1; import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_2; import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_3; +import static com.nvidia.nvcf.util.TestConstants.TEST_VERSION_ID_5; import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatThrownBy; +import com.nvidia.nvcf.persistence.function.entity.FunctionType; import com.nvidia.nvcf.persistence.function.entity.RateLimitUdt; import com.nvidia.nvcf.proto.ClientInvokeRequest; import com.nvidia.nvcf.proto.ClientInvokeResponse; @@ -1088,4 +1090,111 @@ void errorResponseIncludesNcaIdInMetadata() { channel.shutdownNow(); } + Stream llmFunctionVersionArgs() { + return Stream.of( + Arguments.of(TEST_VERSION_ID_1.toString()), + Arguments.of((String) null)); + } + + /** + * LLM workers do not consume the request queues used by classic invocation, so the auth + * call must reject LLM functions up front instead of letting the request time out. + */ + @ParameterizedTest + @MethodSource("llmFunctionVersionArgs") + void rejectsLlmFunctionForClassicInvocation(@Nullable String versionId) { + setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1); + setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.LLM); + var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT, + List.of(SCOPE_INVOKE_FUNCTION), 100); + var ncaIdKey = Metadata.Key.of(TAG_NCA_ID, Metadata.ASCII_STRING_MARSHALLER); + + assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID.toString(), versionId)) + .isInstanceOf(StatusRuntimeException.class) + .satisfies(thrown -> { + var exception = (StatusRuntimeException) thrown; + assertThat(exception.getStatus().getCode()) + .isEqualTo(Status.NOT_FOUND.getCode()); + assertThat(exception.getStatus().getDescription()) + .contains(TEST_FUNCTION_ID.toString()) + .contains("LLM functions cannot be invoked through this endpoint"); + assertThat(exception.getTrailers()).isNotNull(); + assertThat(exception.getTrailers().get(ncaIdKey)).isEqualTo(TEST_NCA_ID); + }); + } + + @Test + void rejectsLlmFunctionForClassicAdminInvocation() { + setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1); + setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.LLM); + var adminAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_ADMIN_SUBJECT, + List.of(ADMIN_SCOPE_INVOKE_FUNCTION), 100); + + assertThatThrownBy(() -> functionAdminAuth(adminAuth, TEST_FUNCTION_ID.toString(), + TEST_VERSION_ID_1.toString())) + .isInstanceOf(StatusRuntimeException.class) + .extracting(thrown -> ((StatusRuntimeException) thrown).getStatus().getCode()) + .isEqualTo(Status.NOT_FOUND.getCode()); + } + + @Test + void llmFunctionWithoutAccessStillReturnsNotFound() { + // TEST_FUNCTION_ID_2 belongs to another account. Access checks must run first so the + // function type of an inaccessible function is not revealed. + setFunctionActive(TEST_FUNCTION_ID_2, TEST_VERSION_ID_2); + setFunctionType(TEST_FUNCTION_ID_2, TEST_VERSION_ID_2, FunctionType.LLM); + var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT, + List.of(SCOPE_INVOKE_FUNCTION), 100); + + assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID_2.toString(), + TEST_VERSION_ID_2.toString())) + .isInstanceOf(StatusRuntimeException.class) + .satisfies(thrown -> { + var status = ((StatusRuntimeException) thrown).getStatus(); + assertThat(status.getCode()).isEqualTo(Status.NOT_FOUND.getCode()); + assertThat(status.getDescription()).doesNotContain("LLM"); + }); + } + + /** + * Families created before type uniformity was enforced may mix LLM and classic versions. + * A versionless classic request must still route to the classic versions. + */ + @Test + void mixedFamilyReturnsOnlyClassicVersions() { + createFunctionVersion(TEST_FUNCTION_ID, TEST_VERSION_ID_5); + setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1); + setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_5); + setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_5, FunctionType.LLM); + var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT, + List.of(SCOPE_INVOKE_FUNCTION), 100); + + var clientInvokeResponse = functionAuth(clientAuth, TEST_FUNCTION_ID.toString(), null); + assertThat(clientInvokeResponse.getFunctionVersionsList()) + .extracting(FunctionVersion::getFunctionVersionId) + .containsExactly(TEST_VERSION_ID_1.toString()); + + // explicitly targeting the LLM version is still rejected + assertThatThrownBy(() -> functionAuth(clientAuth, TEST_FUNCTION_ID.toString(), + TEST_VERSION_ID_5.toString())) + .isInstanceOf(StatusRuntimeException.class) + .extracting(thrown -> ((StatusRuntimeException) thrown).getStatus().getCode()) + .isEqualTo(Status.NOT_FOUND.getCode()); + } + + @Test + void allowsStreamingFunctionForClassicInvocation() { + setFunctionActive(TEST_FUNCTION_ID, TEST_VERSION_ID_1); + setFunctionType(TEST_FUNCTION_ID, TEST_VERSION_ID_1, FunctionType.STREAMING); + var clientAuth = MOCK_OAUTH2_TOKEN_SERVER.getJwt(TEST_CLIENT_SUBJECT, + List.of(SCOPE_INVOKE_FUNCTION), 100); + + var clientInvokeResponse = functionAuth( + clientAuth, TEST_FUNCTION_ID.toString(), TEST_VERSION_ID_1.toString()); + + assertThat(clientInvokeResponse.getFunctionVersionsList()) + .extracting(FunctionVersion::getFunctionVersionId) + .containsExactly(TEST_VERSION_ID_1.toString()); + } + } diff --git a/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/rest/function/invocation/BaseFunctionInvocationTest.java b/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/rest/function/invocation/BaseFunctionInvocationTest.java index bf036d0836..5e1e912d94 100644 --- a/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/rest/function/invocation/BaseFunctionInvocationTest.java +++ b/src/control-plane-services/cloud-functions/nvcf-core/src/test/java/com/nvidia/nvcf/rest/function/invocation/BaseFunctionInvocationTest.java @@ -333,6 +333,12 @@ protected FunctionEntity setFunctionRateLimit(UUID functionId, UUID functionVers return functionsRepository.save(entity); } + protected void createFunctionVersion(UUID functionId, UUID functionVersionId) { + testService.createTestFunctionEntity(functionId, functionVersionId, + TEST_NCA_ID, TEST_FUNCTION_NAME, + FunctionStatus.DEPLOYING); + } + protected FunctionEntity setFunctionActive(UUID functionId, UUID functionVersionId) { return setFunctionStatus(functionId, functionVersionId, FunctionStatus.ACTIVE); } diff --git a/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/mod.rs b/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/mod.rs index 593b44a66f..40a1e8d06a 100644 --- a/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/mod.rs +++ b/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/mod.rs @@ -45,6 +45,10 @@ pub const VERSION_ID_1: Uuid = uuid!("26597542-1782-4a18-aa02-504ba0598202"); pub const VERSION_ID_2: Uuid = uuid!("49331123-201a-407d-9cbc-7bb328fd0295"); pub const VERSION_ID_3: Uuid = uuid!("4117e32b-1b96-4be2-8cd2-9da4047daa05"); pub const VERSION_ID_4: Uuid = uuid!("2bf046f7-37c1-40b8-b2c2-8f8d260f7645"); +#[allow(unused)] +pub const LLM_FUNCTION_ID: Uuid = uuid!("5d0c8a3e-9f4b-4e62-8a1d-3b7e2c9f0a14"); +#[allow(unused)] +pub const LLM_VERSION_ID: Uuid = uuid!("e2a7b4c1-6d3f-4a85-9c02-7f1e8b5d3a96"); pub const LOCALSTACK_REGION: &str = "us-east-1"; pub const ASSETS_BUCKET: &str = "assets-bucket"; pub const RESULTS_BUCKET: &str = "results-bucket"; @@ -141,6 +145,7 @@ async fn mock_nvcf_api() -> ApiMockServer { functions: vec![VERSION_ID_1, VERSION_ID_2], has_rate_limit: false, sync_check: false, + is_llm: false, }, ), // FUNCTION_ID_2_RATELIMIT_SYNC has rate limit and sync check @@ -150,6 +155,7 @@ async fn mock_nvcf_api() -> ApiMockServer { functions: vec![VERSION_ID_3], has_rate_limit: true, sync_check: true, + is_llm: false, }, ), // FUNCTION_ID_3_RATELIMIT_ASYNC has rate limit and async check @@ -159,6 +165,16 @@ async fn mock_nvcf_api() -> ApiMockServer { functions: vec![VERSION_ID_4], has_rate_limit: true, sync_check: false, + is_llm: false, + }, + ), + ( + LLM_FUNCTION_ID, + FunctionMetadata { + functions: vec![LLM_VERSION_ID], + has_rate_limit: false, + sync_check: false, + is_llm: true, }, ), ] diff --git a/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/nvcf_api_mock.rs b/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/nvcf_api_mock.rs index 8b5b85613e..5569a4351f 100644 --- a/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/nvcf_api_mock.rs +++ b/src/invocation-plane-services/http-invocation/crates/server/tests/mocks/nvcf_api_mock.rs @@ -44,6 +44,8 @@ pub struct FunctionMetadata { pub functions: Vec, pub has_rate_limit: bool, pub sync_check: bool, + /// LLM functions are rejected by nvcf-service for classic invocation. + pub is_llm: bool, } #[derive(Debug)] @@ -82,6 +84,15 @@ impl ApiMock { // Include nca_id in error metadata (matching nvcf-service GrpcNcaIdException behavior) Self::status_with_nca_id(tonic::Code::NotFound, "function not found", &client.nca_id) })?; + if function_metadata.is_llm { + return Err(Self::status_with_nca_id( + tonic::Code::NotFound, + &format!( + "Function id '{function_id}': LLM functions cannot be invoked through this endpoint. Use the LLM API instead." + ), + &client.nca_id, + )); + } let requested_function_version_id = request.function_version_id; Ok(Response::new(ClientInvokeResponse { function_id: function_id.into(), diff --git a/src/invocation-plane-services/http-invocation/crates/server/tests/test_error_codes.rs b/src/invocation-plane-services/http-invocation/crates/server/tests/test_error_codes.rs index f990040206..ebfc225ecf 100644 --- a/src/invocation-plane-services/http-invocation/crates/server/tests/test_error_codes.rs +++ b/src/invocation-plane-services/http-invocation/crates/server/tests/test_error_codes.rs @@ -16,19 +16,24 @@ mod mocks; use crate::mocks::nvcf_worker_mock::PublishMode; +use async_nats::jetstream::{context::GetStreamErrorKind, ErrorCode}; use axum::{ body::Body, http::{self, header, HeaderValue, Method, StatusCode}, }; use futures::future; +use http_body_util::BodyExt; use mocks::{ fixtures, nvcf_worker_mock::{ DefaultWorkHandler, DroppableBackgroundWorker, FailedWorkHandler, Worker, WorkerProperties, }, - API_KEY, FUNCTION_ID, INSTANCE_ID, VERSION_ID_1, + API_KEY, FUNCTION_ID, INSTANCE_ID, LLM_FUNCTION_ID, LLM_VERSION_ID, VERSION_ID_1, }; use nvcf_invocation_service::app::app; +use nvcf_invocation_service::nats::NatsService; +use nvcf_invocation_service::settings::GrpcClientConfig; +use problem_details::ProblemDetails; use std::time::{Duration, Instant}; use tokio::{task::JoinHandle, time::timeout}; use tower::{Service, ServiceExt}; @@ -342,3 +347,97 @@ async fn test_delayed_worker_start() -> anyhow::Result<()> { } Ok(()) } + +/// LLM functions are rejected by the auth call, so classic invocation must fail fast with a 404 +/// on every entry point and must not publish a request that no worker will consume. +#[tokio::test] +async fn test_llm_function_rejected_without_publishing() -> anyhow::Result<()> { + let (_localstack, _nats, _mock_nvcf_api, config) = fixtures().await; + let mut app = app(config.clone(), None).await?; + let app = ServiceExt::>::ready(&mut app).await?; + + let requests = [ + http::Request::builder() + .method(Method::POST) + .uri(format!( + "/v2/nvcf/pexec/functions/{LLM_FUNCTION_ID}/versions/{LLM_VERSION_ID}" + )) + .header(header::AUTHORIZATION, format!("Bearer {API_KEY}")) + .body(Body::from("a body"))?, + http::Request::builder() + .method(Method::POST) + .uri(format!("/v2/nvcf/pexec/functions/{LLM_FUNCTION_ID}")) + .header(header::AUTHORIZATION, format!("Bearer {API_KEY}")) + .body(Body::from("a body"))?, + http::Request::builder() + .method(Method::POST) + .uri(format!("/v2/nvcf/exec/functions/{LLM_FUNCTION_ID}")) + .header(header::AUTHORIZATION, format!("Bearer {API_KEY}")) + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(r#"{"requestBody":{"input":"hi"}}"#))?, + // host form, as used by ./v1/responses + http::Request::builder() + .method(Method::POST) + .uri("/v1/responses") + .header( + header::HOST, + format!("{LLM_FUNCTION_ID}.example.nvidia.com"), + ) + .header(header::AUTHORIZATION, format!("Bearer {API_KEY}")) + .header(header::CONTENT_TYPE, "application/json") + .body(Body::from(r#"{"input":"hi"}"#))?, + ]; + + for request in requests { + let uri = request.uri().clone(); + let start = Instant::now(); + let response = timeout(Duration::from_secs(10), app.call(request)) + .await + .unwrap_or_else(|_| panic!("{uri}: request should not wait for a worker"))?; + assert!( + start.elapsed() < Duration::from_secs(5), + "{uri}: took {:?}", + start.elapsed() + ); + assert_eq!(response.status(), StatusCode::NOT_FOUND, "{uri}"); + assert_eq!( + response.headers().get(header::CONTENT_TYPE), + Some(&HeaderValue::from_static("application/problem+json")), + "{uri}" + ); + let body = response.into_body().collect().await?.to_bytes(); + let problem: ProblemDetails = serde_json::from_slice(&body)?; + assert_eq!( + problem.status, + Some(StatusCode::NOT_FOUND), + "{uri}: {problem:?}" + ); + assert!( + problem + .detail + .as_deref() + .is_some_and(|detail| detail.contains("LLM functions cannot be invoked")), + "{uri}: {problem:?}" + ); + } + + // publishing creates the request stream on demand, so its absence proves nothing was queued + let nats_service = NatsService::new( + &config.nats_properties, + "http://dummy.localhost", + None, + &GrpcClientConfig::default(), + ) + .await?; + let stream_name = nats_service.request_stream_name(LLM_VERSION_ID); + match nats_service.jetstream().get_stream(&stream_name).await { + Ok(_) => panic!("no request stream should exist for the rejected LLM function"), + Err(err) => match err.kind() { + GetStreamErrorKind::JetStream(js_err) + if js_err.error_code() == ErrorCode::STREAM_NOT_FOUND => {} + _ => return Err(err.into()), + }, + } + + Ok(()) +}