diff --git a/core/src/main/scala/sttp/model/sse/ServerSentEvent.scala b/core/src/main/scala/sttp/model/sse/ServerSentEvent.scala index 70326cd0..3919a58a 100644 --- a/core/src/main/scala/sttp/model/sse/ServerSentEvent.scala +++ b/core/src/main/scala/sttp/model/sse/ServerSentEvent.scala @@ -9,15 +9,27 @@ case class ServerSentEvent( retry: Option[Int] = None ) { override def toString: String = { - val _data = data.map(_.split("\n")).map(_.map(line => Some(s"data: $line"))).getOrElse(Array.empty[Option[String]]) - val _event = eventType.map(event => s"event: $event") - val _id = id.map(id => s"id: $id") + val _data = data + .map(ServerSentEvent.splitOnLineTerminators) + .map(_.map(line => Some(s"data: $line"))) + .getOrElse(Array.empty[Option[String]]) + val _event = eventType.map(event => s"event: ${ServerSentEvent.removeLineTerminators(event)}") + val _id = id.map(id => s"id: ${ServerSentEvent.removeLineTerminators(id)}") val _retry = retry.map(retryCount => s"retry: $retryCount") (_data :+ _event :+ _id :+ _retry).flatten.mkString("\n") } } object ServerSentEvent { + private val LineTerminators = "\r\n|\r|\n" + + // performance: split("\n") skips the regex engine; with no CR, LF is the only terminator, so it's equivalent + private def splitOnLineTerminators(s: String): Array[String] = + if (s.indexOf('\r') < 0) s.split("\n", -1) else s.split(LineTerminators, -1) + + private def removeLineTerminators(s: String): String = + if (s.indexOf('\r') < 0 && s.indexOf('\n') < 0) s else s.replaceAll(LineTerminators, "") + // https://html.spec.whatwg.org/multipage/server-sent-events.html def parse(event: List[String]): ServerSentEvent = { event.foldLeft(ServerSentEvent()) { (event, line) => diff --git a/core/src/test/scala/sttp/model/sse/ServerSentEventTest.scala b/core/src/test/scala/sttp/model/sse/ServerSentEventTest.scala index 5071c641..dc36118a 100644 --- a/core/src/test/scala/sttp/model/sse/ServerSentEventTest.scala +++ b/core/src/test/scala/sttp/model/sse/ServerSentEventTest.scala @@ -64,4 +64,48 @@ class ServerSentEventTest extends AnyFlatSpec with Matchers { |data: some data info 2 |data: some data info 3""".stripMargin } + + "composeSSE" should "split data on all line terminators" in { + val sse = ServerSentEvent(Some("line 1\r\nline 2\rline 3\nline 4")) + + sse.toString shouldBe + s"""data: line 1 + |data: line 2 + |data: line 3 + |data: line 4""".stripMargin + } + + "composeSSE" should "remove line terminators from the event type" in { + val sse = ServerSentEvent(eventType = Some("a\ndata: injected\rb\r\nc")) + sse.toString shouldBe "event: adata: injectedbc" + } + + "composeSSE" should "remove line terminators from the id" in { + val sse = ServerSentEvent(id = Some("a\ndata: injected\rb\r\nc")) + sse.toString shouldBe "id: adata: injectedbc" + } + + "composeSSE" should "not allow injecting fields through data, the event type or the id" in { + val malicious = "x\r\nevent: injected\rid: injected\ndata: injected" + val sse = ServerSentEvent(Some(malicious), Some(malicious), Some(malicious), Some(10)) + + ServerSentEvent.parse(sse.toString.split("\n").toList) shouldBe ServerSentEvent( + Some("x\nevent: injected\nid: injected\ndata: injected"), + Some("xevent: injectedid: injecteddata: injected"), + Some("xevent: injectedid: injecteddata: injected"), + Some(10) + ) + } + + "composeSSE" should "keep a trailing line terminator in data" in { + ServerSentEvent(Some("a\n")).toString shouldBe "data: a\ndata: " + ServerSentEvent(Some("a\r")).toString shouldBe "data: a\ndata: " + ServerSentEvent(Some("a\r\n")).toString shouldBe "data: a\ndata: " + ServerSentEvent(Some("\n")).toString shouldBe "data: \ndata: " + } + + "composeSSE" should "round-trip data with a trailing line terminator" in { + val sse = ServerSentEvent(Some("a\n")) + ServerSentEvent.parse(sse.toString.split("\n").toList) shouldBe sse + } }