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
Expand Up @@ -2,21 +2,28 @@

import edu.umd.cs.findbugs.annotations.SuppressFBWarnings;
import java.util.function.Function;

import javax.annotation.Nonnull;

/**
* Map with weakly referenced keys.
*
* <p>Keys must not be {@code null}: the backing {@code WeakConcurrentMap} rejects null keys and
* throws.
*/
public interface WeakMap<K, V> {
int size();

boolean containsKey(K target);
boolean containsKey(@Nonnull K target);

V get(K key);
V get(@Nonnull K key);

void put(K key, V value);
void put(@Nonnull K key, V value);

void putIfAbsent(K key, V value);
void putIfAbsent(@Nonnull K key, V value);

V computeIfAbsent(K key, Function<? super K, ? extends V> supplier);
V computeIfAbsent(@Nonnull K key, Function<? super K, ? extends V> supplier);

V remove(K key);
V remove(@Nonnull K key);

abstract class Supplier {
private static volatile Supplier SUPPLIER;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,11 @@ public static void onExit(
return;
}
final ServletContext context = request.getServletContext();
if (context == null) {
// some request copies (e.g. Wicket or Atmosphere websocket requests) have no servlet
// context
return;
}

if (InstrumentationContext.get(ServletContext.class, SessionTrackingMode.class).get(context)
!= null) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -460,6 +460,24 @@ class JakartaHttpServletRequestInstrumentationTest extends InstrumentationSpecif
suite << testSuite()
}

void 'getSession tolerates a null servlet context'() {
setup:
final module = Mock(ApplicationModule)
InstrumentationBridge.registerIastModule(module)
final session = Mock(HttpSession)
final delegate = Mock(HttpServletRequest)
final request = new CustomRequest(request: delegate)

when:
final result = request.getSession()

then:
result.is(session)
1 * delegate.getSession() >> session
1 * delegate.getServletContext() >> null
0 * module._
}

protected <E> E runUnderIastTrace(Closure<E> cl) {
final ddctx = new TagContext().withRequestContextDataIast(iastCtx)
final span = TEST_TRACER.startSpan("test", "test-iast-span", ddctx)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,11 @@ public static void onExit(
return;
}
final ServletContext context = request.getServletContext();
if (context == null) {
// some request copies (e.g. Wicket or Atmosphere websocket requests) have no servlet
// context
return;
}
if (InstrumentationContext.get(ServletContext.class, SessionTrackingMode.class).get(context)
!= null) {
return;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ import javax.servlet.annotation.WebServlet
import javax.servlet.http.HttpServlet
import javax.servlet.http.HttpServletRequest
import javax.servlet.http.HttpServletResponse
import javax.servlet.http.HttpSession

import static datadog.trace.agent.test.base.HttpServerTest.ServerEndpoint.CUSTOM_EXCEPTION
import static datadog.trace.agent.test.base.HttpServerTest.ServerEndpoint.ERROR
Expand Down Expand Up @@ -570,6 +571,32 @@ class JettyServlet3ServeFromAsyncTimeout extends JettyServlet3Test {

class IastJettyServlet3ForkedTest extends JettyServlet3TestSync {

void 'getSession tolerates a null servlet context'() {
setup:
final module = Mock(ApplicationModule)
InstrumentationBridge.registerIastModule(module)
final session = Mock(HttpSession)
final delegate = Mock(HttpServletRequest)
final request = new CustomRequest(request: delegate)

when:
final result = request.getSession()

then:
result.is(session)
1 * delegate.getSession() >> session
1 * delegate.getServletContext() >> null
0 * module._

cleanup:
InstrumentationBridge.clearIastModules()
}

private static class CustomRequest implements HttpServletRequest {
@Delegate
private HttpServletRequest request
}

@Override
Class<Servlet> servlet() {
return TestServlet3.GetSession
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package datadog.trace.bootstrap;

import java.util.function.Function;
import javax.annotation.Nonnull;
import javax.annotation.Nullable;

/**
Expand All @@ -9,6 +10,9 @@
* <p>Context instances are weakly referenced and will be garbage collected when their corresponding
* key instance is collected.
*
* <p>Keys must not be {@code null}: the backing weak maps reject null keys and throw. Callers that
* take a key from application code (for example a servlet request's context) must check it first.
*
* @param <K> key type to do context lookups
* @param <C> context type
*/
Expand Down Expand Up @@ -38,15 +42,15 @@ default C apply(Object key) {
* @return context instance; {@code null} if the key had no context
*/
@Nullable
C get(K key);
C get(@Nonnull K key);

/**
* Unconditionally put new context instance for the given key.
*
* @param key the context key
* @param context context instance to save
*/
void put(K key, C context);
void put(@Nonnull K key, C context);

/**
* Gets the context instance for the given key. If no context exists then associate it with the
Expand All @@ -56,7 +60,7 @@ default C apply(Object key) {
* @param context new context instance
* @return existing context instance if present; otherwise new instance
*/
C getOrPut(K key, C context);
C getOrPut(@Nonnull K key, C context);

/**
* Gets the context instance for the given key. If no context exists then create one using the
Expand All @@ -66,7 +70,7 @@ default C apply(Object key) {
* @param contextFactory factory instance to produce new context instances
* @return existing context instance if present; otherwise new instance
*/
default C getOrCreate(K key, Factory<C> contextFactory) {
default C getOrCreate(@Nonnull K key, Factory<C> contextFactory) {
return getOrCompute(key, contextFactory);
}

Expand All @@ -78,7 +82,7 @@ default C getOrCreate(K key, Factory<C> contextFactory) {
* @param contextFactory factory instance to produce new context instances
* @return existing context instance if present; otherwise new instance
*/
C getOrCompute(K key, Function<? super K, C> contextFactory);
C getOrCompute(@Nonnull K key, Function<? super K, C> contextFactory);

/**
* Removes the context instance for the given key.
Expand All @@ -87,5 +91,5 @@ default C getOrCreate(K key, Factory<C> contextFactory) {
* @return removed context instance; {@code null} if the key had no context
*/
@Nullable
C remove(K key);
C remove(@Nonnull K key);
}
Loading