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
92 changes: 92 additions & 0 deletions app/src/test/kotlin/com/refinvest/RefinvestApplicationTests.kt
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,98 @@ class RefinvestApplicationTests(
assertTrue(response.body().contains("\"versions\":[]"), response.body())
}

@Test
fun `defines a strategy version through HTTP and returns it from strategy retrieval`() {
val created = HttpClient.newHttpClient().send(
HttpRequest.newBuilder(URI("http://localhost:$port/strategies"))
.header("Content-Type", "application/json")
.POST(HttpRequest.BodyPublishers.ofString("{\"name\":\"volatility hypothesis\"}"))
.build(),
HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8),
)
val strategyId = "\"id\":\"(\\d+)\"".toRegex().find(created.body())?.groupValues?.get(1)
assertTrue(created.statusCode() == 201 && strategyId != null, created.body())

val defined = HttpClient.newHttpClient().send(
HttpRequest.newBuilder(URI("http://localhost:$port/strategies/$strategyId/versions"))
.header("Content-Type", "application/json")
.POST(
HttpRequest.BodyPublishers.ofString(
"""{
| "primarySignalAsset":"QQQ",
| "conditions":[{
| "operator":"LT",
| "operandA":{"asset":"QQQ","metric":"RETURN","window":5},
| "operandB":-0.07
| }],
| "executionAsset":"TQQQ",
| "lag":3,
| "exit":{"holdingSignalSessions":5}
|}""".trimMargin(),
),
)
.build(),
HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8),
)

assertTrue(defined.statusCode() == 201, defined.body())
assertTrue(defined.body().contains("\"strategyId\":\"$strategyId\""), defined.body())
assertTrue(jdbcTemplate.queryForObject("select count(*) from strategy_versions", Long::class.java) == 1L)

val retrieved = HttpClient.newHttpClient().send(
HttpRequest.newBuilder(URI("http://localhost:$port/strategies/$strategyId")).GET().build(),
HttpResponse.BodyHandlers.ofString(),
)

assertTrue(retrieved.statusCode() == 200, retrieved.body())
assertTrue(retrieved.body().contains("\"primarySignalAsset\":\"QQQ\""), retrieved.body())
assertTrue(retrieved.body().contains("\"latestVersionId\":"), retrieved.body())
}

@Test
fun `rejects VIX as an execution asset when defining a strategy version`() {
val created = HttpClient.newHttpClient().send(
HttpRequest.newBuilder(URI("http://localhost:$port/strategies"))
.header("Content-Type", "application/json")
.POST(HttpRequest.BodyPublishers.ofString("{\"name\":\"volatility hypothesis\"}"))
.build(),
HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8),
)
val strategyId = "\"id\":\"(\\d+)\"".toRegex().find(created.body())?.groupValues?.get(1)
assertTrue(created.statusCode() == 201 && strategyId != null, created.body())

val response = HttpClient.newHttpClient().send(
HttpRequest.newBuilder(URI("http://localhost:$port/strategies/$strategyId/versions"))
.header("Content-Type", "application/json")
.POST(
HttpRequest.BodyPublishers.ofString(
"""{
| "primarySignalAsset":"QQQ",
| "conditions":[{
| "operator":"GT",
| "operandA":{"asset":"QQQ","metric":"SIMPLE"},
| "operandB":1
| }],
| "executionAsset":"VIX",
| "lag":0,
| "exit":{"holdingSignalSessions":1}
|}""".trimMargin(),
),
)
.build(),
HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8),
)

assertTrue(response.statusCode() == 400, response.body())
assertTrue(
jdbcTemplate.queryForObject(
"select count(*) from strategy_versions where strategy_id = ?",
Long::class.java,
strategyId.toLong(),
) == 0L,
)
}

@Test
fun `returns not found for an unknown strategy`() {
val response = HttpClient.newHttpClient().send(
Expand Down
1 change: 1 addition & 0 deletions gradle/libs.versions.toml
Original file line number Diff line number Diff line change
Expand Up @@ -23,3 +23,4 @@ postgresql = { module = "org.postgresql:postgresql" }
h2 = { module = "com.h2database:h2" }
jackson-module-kotlin = { module = "tools.jackson.module:jackson-module-kotlin" }
spring-context = { module = "org.springframework:spring-context", version.ref = "spring-framework" }
spring-tx = { module = "org.springframework:spring-tx", version.ref = "spring-framework" }
Original file line number Diff line number Diff line change
@@ -0,0 +1,43 @@
package com.refinvest.core.strategy.adapter.out.persistence

import jakarta.persistence.Column
import jakarta.persistence.Entity
import jakarta.persistence.GeneratedValue
import jakarta.persistence.GenerationType
import jakarta.persistence.Id
import jakarta.persistence.JoinColumn
import jakarta.persistence.ManyToOne
import jakarta.persistence.Table

@Entity
@Table(name = "strategy_version_conditions")
class ConditionJpaEntity(
@Id
@GeneratedValue(strategy = GenerationType.IDENTITY)
var id: Long? = null,
@ManyToOne
@JoinColumn(name = "strategy_version_id", nullable = false)
var strategyVersion: StrategyVersionJpaEntity,
@Column(name = "condition_order", nullable = false)
var conditionOrder: Int,
@Column(nullable = false)
var operator: String,
@Column(name = "logical_combinator")
var logicalCombinator: String?,
@Column(name = "operand_a_asset", nullable = false)
var operandAAsset: String,
@Column(name = "operand_a_metric", nullable = false)
var operandAMetric: String,
@Column(name = "operand_a_window")
var operandAWindow: Int?,
@Column(name = "operand_b_kind", nullable = false)
var operandBKind: String,
@Column(name = "operand_b_literal")
var operandBLiteral: Double?,
@Column(name = "operand_b_asset")
var operandBAsset: String?,
@Column(name = "operand_b_metric")
var operandBMetric: String?,
@Column(name = "operand_b_window")
var operandBWindow: Int?,
)
Original file line number Diff line number Diff line change
@@ -1,8 +1,20 @@
package com.refinvest.core.strategy.adapter.out.persistence

import com.refinvest.core.strategy.domain.StrategyId
import com.refinvest.core.strategy.domain.StrategyVersionId
import com.refinvest.core.strategy.domain.AssetSymbol
import com.refinvest.core.strategy.domain.ComparisonOperator
import com.refinvest.core.strategy.domain.Condition
import com.refinvest.core.strategy.domain.LiteralValue
import com.refinvest.core.strategy.domain.LogicalCombinator
import com.refinvest.core.strategy.domain.MetricOperand
import com.refinvest.core.strategy.domain.MetricReference
import com.refinvest.core.strategy.domain.MetricType
import com.refinvest.core.strategy.domain.SignalSessions
import com.refinvest.core.strategy.domain.TimeBasedExit
import com.refinvest.core.strategy.port.outbound.StrategyReadModel
import com.refinvest.core.strategy.port.outbound.StrategyReader
import com.refinvest.core.strategy.port.outbound.StrategyVersionReadModel
import org.springframework.stereotype.Repository

@Repository
Expand All @@ -11,11 +23,45 @@ class JpaStrategyReaderAdapter(
) : StrategyReader {
override fun findById(id: StrategyId): StrategyReadModel? =
strategyJpaReader.findById(id.value)?.let { strategy ->
val versions = strategy.versions.map { it.toReadModel() }
StrategyReadModel(
id = StrategyId(strategy.id),
name = strategy.name,
createdAt = strategy.createdAt,
latestVersionId = null,
latestVersionId = versions.lastOrNull()?.id,
versions = versions,
)
}

private fun StrategyVersionJpaEntity.toReadModel(): StrategyVersionReadModel = StrategyVersionReadModel(
id = StrategyVersionId(id),
createdAt = createdAt,
primarySignalAsset = AssetSymbol.valueOf(primarySignalAsset),
conditions = conditions.map { it.toDomain() },
executionAsset = AssetSymbol.valueOf(executionAsset),
lag = SignalSessions(lag),
exit = TimeBasedExit(holdingSignalSessions),
)

private fun ConditionJpaEntity.toDomain(): Condition = Condition(
operator = ComparisonOperator.valueOf(operator),
logicalCombinator = logicalCombinator?.let(LogicalCombinator::valueOf),
operandA = MetricReference(AssetSymbol.valueOf(operandAAsset), MetricType.valueOf(operandAMetric), operandAWindow),
operandB = when (operandBKind) {
LITERAL -> LiteralValue(requireNotNull(operandBLiteral))
METRIC -> MetricOperand(
MetricReference(
AssetSymbol.valueOf(requireNotNull(operandBAsset)),
MetricType.valueOf(requireNotNull(operandBMetric)),
operandBWindow,
),
)
else -> error("Unsupported operand B kind: $operandBKind")
},
)

private companion object {
const val LITERAL = "LITERAL"
const val METRIC = "METRIC"
}
}
Original file line number Diff line number Diff line change
@@ -1,21 +1,128 @@
package com.refinvest.core.strategy.adapter.out.persistence

import com.refinvest.core.strategy.domain.Strategy
import com.refinvest.core.strategy.domain.StrategyId
import com.refinvest.core.strategy.domain.StrategyVersion
import com.refinvest.core.strategy.domain.AssetSymbol
import com.refinvest.core.strategy.domain.ComparisonOperator
import com.refinvest.core.strategy.domain.Condition
import com.refinvest.core.strategy.domain.ConditionOperand
import com.refinvest.core.strategy.domain.LiteralValue
import com.refinvest.core.strategy.domain.LogicalCombinator
import com.refinvest.core.strategy.domain.MemberId
import com.refinvest.core.strategy.domain.MetricOperand
import com.refinvest.core.strategy.domain.MetricReference
import com.refinvest.core.strategy.domain.MetricType
import com.refinvest.core.strategy.domain.SignalSessions
import com.refinvest.core.strategy.domain.StrategyVersionId
import com.refinvest.core.strategy.domain.TimeBasedExit
import com.refinvest.core.strategy.port.outbound.StrategyStore
import org.springframework.stereotype.Repository

@Repository
class JpaStrategyStoreAdapter(
private val strategyJpaStore: StrategyJpaStore,
) : StrategyStore {
override fun findById(id: StrategyId): Strategy? =
strategyJpaStore.findById(id.value).orElse(null)?.toDomain()

override fun save(strategy: Strategy) {
strategyJpaStore.save(
StrategyJpaEntity(
id = strategy.id.value,
memberId = strategy.memberId.value,
name = strategy.name,
createdAt = strategy.createdAt,
),
val entity = StrategyJpaEntity(
id = strategy.id.value,
memberId = strategy.memberId.value,
name = strategy.name,
createdAt = strategy.createdAt,
)
entity.replaceVersions(strategy.versions.map { it.toEntity(entity) })
strategyJpaStore.save(entity)
}

private fun StrategyJpaEntity.toDomain(): Strategy = Strategy.create(
id = StrategyId(id),
memberId = MemberId(memberId),
name = name,
createdAt = createdAt,
).also { strategy -> versions.forEach { strategy.addVersion(it.toDomain()) } }

private fun StrategyVersionJpaEntity.toDomain(): StrategyVersion = StrategyVersion.create(
id = StrategyVersionId(id),
strategyId = StrategyId(strategy.id),
createdAt = createdAt,
primarySignalAsset = AssetSymbol.valueOf(primarySignalAsset),
conditions = conditions.map { it.toDomain() },
executionAsset = AssetSymbol.valueOf(executionAsset),
lag = SignalSessions(lag),
exit = TimeBasedExit(holdingSignalSessions),
)

private fun StrategyVersion.toEntity(strategy: StrategyJpaEntity): StrategyVersionJpaEntity {
val entity = StrategyVersionJpaEntity(
id = id.value,
strategy = strategy,
createdAt = createdAt,
primarySignalAsset = primarySignalAsset.name,
executionAsset = executionAsset.name,
lag = lag.value,
holdingSignalSessions = exit.holdingSignalSessions,
)
entity.conditions += conditions.mapIndexed { index, condition -> condition.toEntity(entity, index) }
return entity
}

private fun ConditionJpaEntity.toDomain(): Condition = Condition(
operator = ComparisonOperator.valueOf(operator),
logicalCombinator = logicalCombinator?.let(LogicalCombinator::valueOf),
operandA = MetricReference(AssetSymbol.valueOf(operandAAsset), MetricType.valueOf(operandAMetric), operandAWindow),
operandB = when (operandBKind) {
LITERAL -> LiteralValue(requireNotNull(operandBLiteral))
METRIC -> MetricOperand(
MetricReference(
AssetSymbol.valueOf(requireNotNull(operandBAsset)),
MetricType.valueOf(requireNotNull(operandBMetric)),
operandBWindow,
),
)
else -> error("Unsupported operand B kind: $operandBKind")
},
)

private fun Condition.toEntity(
strategyVersion: StrategyVersionJpaEntity,
conditionOrder: Int,
): ConditionJpaEntity =
when (val operandB = operandB) {
is LiteralValue -> ConditionJpaEntity(
strategyVersion = strategyVersion,
conditionOrder = conditionOrder,
operator = operator.name,
logicalCombinator = logicalCombinator?.name,
operandAAsset = operandA.asset.name,
operandAMetric = operandA.metric.name,
operandAWindow = operandA.window,
operandBKind = LITERAL,
operandBLiteral = operandB.value,
operandBAsset = null,
operandBMetric = null,
operandBWindow = null,
)
is MetricOperand -> ConditionJpaEntity(
strategyVersion = strategyVersion,
conditionOrder = conditionOrder,
operator = operator.name,
logicalCombinator = logicalCombinator?.name,
operandAAsset = operandA.asset.name,
operandAMetric = operandA.metric.name,
operandAWindow = operandA.window,
operandBKind = METRIC,
operandBLiteral = null,
operandBAsset = operandB.reference.asset.name,
operandBMetric = operandB.reference.metric.name,
operandBWindow = operandB.reference.window,
)
}

private companion object {
const val LITERAL = "LITERAL"
const val METRIC = "METRIC"
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,11 @@ package com.refinvest.core.strategy.adapter.out.persistence

import jakarta.persistence.Column
import jakarta.persistence.Entity
import jakarta.persistence.FetchType
import jakarta.persistence.Id
import jakarta.persistence.OneToMany
import jakarta.persistence.Table
import jakarta.persistence.CascadeType
import java.time.Instant

@Entity
Expand All @@ -17,4 +20,11 @@ class StrategyJpaEntity(
var name: String,
@Column(name = "created_at", nullable = false)
var createdAt: Instant,
)
@OneToMany(mappedBy = "strategy", cascade = [CascadeType.ALL], orphanRemoval = true, fetch = FetchType.EAGER)
var versions: MutableList<StrategyVersionJpaEntity> = mutableListOf(),
) {
fun replaceVersions(newVersions: List<StrategyVersionJpaEntity>) {
versions.clear()
versions += newVersions
}
}
Loading